feat: causal lab

This commit is contained in:
OpenSquared
2026-06-28 18:08:09 +02:00
parent c286c7c000
commit 6ebbf4326e
5 changed files with 481 additions and 133 deletions

View File

@@ -431,6 +431,8 @@ def patch_template(template_id: int, body: dict):
sets.append("instruments=?"); params.append(json.dumps(body["instruments"]))
if "graph_json" in body:
sets.append("graph_json=?"); params.append(json.dumps(body["graph_json"]))
if "calibration_json" in body:
sets.append("calibration_json=?"); params.append(json.dumps(body["calibration_json"]))
if not sets:
conn.close(); return {"ok": True}
@@ -1235,3 +1237,100 @@ def get_calibration():
except Exception as e:
logger.error(f"[causal_lab] calibration: {e}")
raise HTTPException(500, str(e))
@router.post("/api/causal-lab/template/{template_id}/generate-theory")
def generate_theory(template_id: int):
"""GPT-4o-mini génère absorption_days, decay_type et confidence pour le template parent."""
try:
from services.database import get_conn, get_config
from services.causal_graphs import get_template
conn = get_conn()
_init(conn)
tmpl = get_template(conn, template_id)
if not tmpl:
conn.close(); raise HTTPException(404, "Template introuvable")
stats = conn.execute("""
SELECT COUNT(*) as n, AVG(activation_score) as avg_act
FROM causal_event_analyses WHERE template_id = ?
""", (template_id,)).fetchone()
conn.close()
n_analyses = stats["n"] if stats else 0
avg_act = round((stats["avg_act"] or 0.0), 2) if stats else 0.0
graph = tmpl.get("graph_json", {})
coefs = {k: v.get("value") for k, v in graph.get("coefficients", {}).items()}
key = get_config("openai_api_key") or ""
if not key:
raise HTTPException(400, "Clé OpenAI manquante dans la configuration")
import openai
prompt = (
f'Tu es un expert en microstructure de marché et dynamique d\'absorption des chocs de prix.\n\n'
f'Template causal : "{tmpl["name"]}"\n'
f'Catégorie : {tmpl["category"]} / {tmpl.get("sub_type", "")}\n'
f'Description : {tmpl.get("description", "")}\n'
f'Instruments : {tmpl.get("instruments", [])}\n'
f'Coefficients : {json.dumps(coefs, ensure_ascii=False)}\n'
f'Analyses historiques : {n_analyses} (activation directionnelle moy. : {avg_act:.0%})\n\n'
f'Propose les paramètres d\'absorption de l\'impact de marché :\n'
f'- absorption_days : jours calendaires avant absorption à >90% (entier 1-60)\n'
f'- decay_type : "step" (tout-ou-rien), "linear" (déclin linéaire), "exp" (exponentiel)\n'
f'- confidence : confiance 0.0-1.0\n'
f'- rationale : justification courte (max 120 chars)\n\n'
f'Références : décisions taux→3-7j/exp ; CPI/NFP→2-5j/exp ; '
f'géopolitique→5-21j/linear ; PMI secondaire→1-2j/step\n\n'
f'JSON uniquement : {{"absorption_days": N, "decay_type": "...", "confidence": 0.X, "rationale": "..."}}'
)
client = openai.OpenAI(api_key=key)
resp = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{"role": "user", "content": prompt}],
response_format={"type": "json_object"},
temperature=0.2,
max_tokens=200,
)
raw = resp.choices[0].message.content or "{}"
params = json.loads(raw)
absorption_days = max(1, min(60, int(params.get("absorption_days", 7))))
decay_type = params.get("decay_type", "exp")
if decay_type not in ("step", "linear", "exp"):
decay_type = "exp"
confidence = round(max(0.0, min(1.0, float(params.get("confidence", 0.5)))), 2)
rationale = str(params.get("rationale", ""))[:150]
conn2 = get_conn()
existing_calib = dict(tmpl.get("calibration_json") or {})
existing_calib.update({
"absorption_days": absorption_days,
"decay_type": decay_type,
"confidence": confidence,
"theory_rationale": rationale,
"theory_generated_at": datetime.utcnow().isoformat() + "Z",
})
conn2.execute(
"UPDATE causal_graph_templates SET calibration_json=?, updated_at=datetime('now') WHERE id=?",
(json.dumps(existing_calib), template_id),
)
conn2.commit()
conn2.close()
return {
"template_id": template_id,
"absorption_days": absorption_days,
"decay_type": decay_type,
"confidence": confidence,
"rationale": rationale,
}
except HTTPException:
raise
except Exception as e:
logger.error(f"[causal_lab] generate_theory {template_id}: {e}")
raise HTTPException(500, str(e))