feat:: causal lab
This commit is contained in:
@@ -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):
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user