feat:: causal lab

This commit is contained in:
OpenSquared
2026-06-28 15:38:19 +02:00
parent 863ba67610
commit 1e44557551
7 changed files with 321 additions and 84 deletions

View File

@@ -321,6 +321,11 @@ class RecommendRequest(BaseModel):
market_event_id: int
class InstantiateRequest(BaseModel):
market_event_id: int
template_id: int
class AnalyzeRequest(BaseModel):
market_event_id: int
template_id: int
@@ -635,6 +640,97 @@ def recommend_template(body: RecommendRequest):
raise HTTPException(500, str(e))
@router.post("/api/causal-lab/instantiate")
def instantiate_template(body: InstantiateRequest):
"""
GPT-4o-mini suggère les valeurs des nœuds user_input du template
en fonction des spécificités de l'événement (surprise +/-, magnitude, catégorie…).
Retourne { inputs: { node_id: value }, rationale: str }
"""
try:
from services.database import get_conn
conn = get_conn()
_init(conn)
ev_row = conn.execute("SELECT * FROM market_events WHERE id = ?", (body.market_event_id,)).fetchone()
tmpl_row = conn.execute("SELECT * FROM causal_graph_templates WHERE id = ?", (body.template_id,)).fetchone()
conn.close()
if not ev_row or not tmpl_row:
raise HTTPException(404, "event ou template introuvable")
event = dict(ev_row)
template = dict(tmpl_row)
graph_json = json.loads(template.get("graph_json") or "{}")
input_mapping = graph_json.get("input_mapping", {})
# Only user_input nodes need to be suggested (auto-sourced nodes are pulled from DB)
user_nodes = {
k: v for k, v in input_mapping.items()
if not v.get("source") or v["source"] == "user_input"
}
if not user_nodes:
return {"inputs": {}, "rationale": "Aucun nœud manuel — tout est auto-récupéré depuis la BD"}
prompt = f"""Tu es un analyste financier. Tu dois configurer un graphe causal pour un événement de marché précis.
Événement : {event.get('name')}
Catégorie : {event.get('category')} / {event.get('sub_type', '')}
Date : {event.get('start_date')}
Description : {(event.get('description') or '')[:400]}
Score impact : {event.get('impact_score', 0.5)}
Valeur réelle : {event.get('actual_value')} | Valeur attendue : {event.get('expected_value')} | Surprise % : {event.get('surprise_pct')}
Template : {template.get('name')} (catégorie : {template.get('category')})
Nœuds à configurer (entrées manuelles) :
{json.dumps(user_nodes, indent=2, ensure_ascii=False)}
Règles :
- Surprise NÉGATIVE → valeurs négatives pour les nœuds de surprise
- Surprise POSITIVE → valeurs positives
- Utilise les plages "range" si fournies (ex: [-3, 3] pour z-score)
- Magnitude proportionnelle à l'ampleur de la surprise
- Si surprise_pct disponible, utilise-la pour calibrer (ex: -29.1% → nœud_surprise ≈ -2.5)
Retourne UNIQUEMENT un JSON : {{ "node_id": valeur_numerique, ... }}
Inclure uniquement les nœuds listés ci-dessus."""
client = _openai_client()
resp = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{"role": "user", "content": prompt}],
response_format={"type": "json_object"},
temperature=0.2,
max_tokens=400,
)
raw = resp.choices[0].message.content or "{}"
suggested = json.loads(raw)
# Validate: only keep known nodes with numeric values
validated: dict = {}
for k, v in suggested.items():
if k in user_nodes:
try:
validated[k] = float(v)
except (TypeError, ValueError):
pass
surprise = event.get("surprise_pct")
rationale = (
f"Surprise {'+' if (surprise or 0) >= 0 else ''}{surprise:.1f}%" if surprise is not None
else "Calibration basée sur la description de l'événement"
)
return {"inputs": validated, "rationale": rationale}
except HTTPException:
raise
except Exception as e:
logger.error(f"[causal_lab] instantiate: {e}")
raise HTTPException(500, str(e))
@router.post("/api/causal-lab/analyze")
def analyze_event(body: AnalyzeRequest):
"""