File size: 4,831 Bytes
db4ba8d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
"""
TradeFlow AI — CEISA Poll Tasks (T-060) + Learning Agent support (T-065, T-066)
"""
from __future__ import annotations

import asyncio
import json
import logging

logger = logging.getLogger("tasks.ceisa_poll")


def poll_active_batches() -> dict:
    """
    Celery Beat (every 30s): poll CEISA for all SUBMITTED batches.
    Resumes LangGraph graph from poll_status node.
    """
    try:
        return asyncio.run(_poll_async())
    except Exception as e:
        logger.error(f"CEISA poll task error: {e}")
        return {"polled": 0, "error": str(e)}


async def _poll_async() -> dict:
    from packages.agents.src.graph import get_graph  # type: ignore

    from ..db.database import get_async_session  # type: ignore

    graph = get_graph()
    polled = 0

    async with get_async_session() as db:
        rows = await db.execute(
            "SELECT id, langgraph_thread_id, ceisa_aju FROM batches "
            "WHERE status = 'SUBMITTED' AND ceisa_aju IS NOT NULL LIMIT 50"
        )
        batches = [dict(r) for r in (rows or [])]

    for batch in batches:
        try:
            config = {"configurable": {"thread_id": batch["langgraph_thread_id"]}}
            await graph.ainvoke(None, config=config)
            polled += 1
        except Exception as e:
            logger.error(f"Poll failed for batch {batch['id']}: {e}")

    return {"polled": polled, "total": len(batches)}


async def record_learning_outcome(
    batch_id: str,
    approved: bool,
    ceisa_status: str,
    corrections: list[dict],
    crs_score: float | None = None,
    rejection_probability: float | None = None,
) -> None:
    """Persist outcome to learning_outcomes for adaptive retraining."""
    from ..db.database import get_async_session  # type: ignore

    async with get_async_session() as db:
        await db.execute(
            """
            INSERT INTO learning_outcomes
              (batch_id, approved, ceisa_status, correction_count,
               corrections_json, crs_score, rejection_probability, created_at)
            VALUES ($1,$2,$3,$4,$5,$6,$7,NOW())
            ON CONFLICT (batch_id) DO UPDATE
              SET ceisa_status=EXCLUDED.ceisa_status, approved=EXCLUDED.approved
            """,
            batch_id, approved, ceisa_status,
            len(corrections), json.dumps(corrections),
            crs_score, rejection_probability,
        )


def check_model_drift() -> dict:
    """Celery beat daily: trigger XGBoost retrain if 100+ new samples."""
    try:
        return asyncio.run(_check_drift_async())
    except Exception as e:
        logger.error(f"Drift check failed: {e}")
        return {"triggered_retrain": False}


async def _check_drift_async() -> dict:
    from ..db.database import get_async_session  # type: ignore
    async with get_async_session() as db:
        row = await db.fetchrow(
            "SELECT COUNT(*) AS n FROM learning_outcomes WHERE created_at > NOW() - INTERVAL '24 hours'"
        )
        new_samples = int(row["n"]) if row else 0

    if new_samples >= 100:
        from .celery_app import celery_app  # type: ignore
        celery_app.send_task("src.tasks.ceisa_poll_tasks.retrain_xgboost")
        return {"triggered_retrain": True, "new_samples": new_samples}
    return {"triggered_retrain": False, "new_samples": new_samples}


def retrain_xgboost() -> dict:
    """Retrain XGBoost predictor from labeled outcomes."""
    try:
        return asyncio.run(_retrain_async())
    except Exception as e:
        logger.error(f"XGBoost retrain error: {e}")
        return {"success": False}


async def _retrain_async() -> dict:
    import os

    import numpy as np
    import xgboost as xgb

    from ..config import settings  # type: ignore
    from ..db.database import get_async_session  # type: ignore

    async with get_async_session() as db:
        rows = await db.fetch(
            "SELECT approved, correction_count, crs_score FROM learning_outcomes "
            "WHERE crs_score IS NOT NULL ORDER BY created_at DESC LIMIT 5000"
        )

    data = [dict(r) for r in (rows or [])]
    if len(data) < settings.XGB_MIN_SAMPLES_FOR_MODEL:
        return {"success": False, "reason": "insufficient_samples", "count": len(data)}

    X = np.array([[r["crs_score"], r["correction_count"]] for r in data], dtype=np.float32)
    y = np.array([0 if r["approved"] else 1 for r in data], dtype=np.float32)
    dtrain = xgb.DMatrix(X, label=y)
    model = xgb.train({"objective": "binary:logistic", "eta": 0.1, "max_depth": 4}, dtrain, 100)

    path = getattr(settings, "XGBOOST_MODEL_PATH", "models/rejection_predictor.json")
    os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
    model.save_model(path)
    logger.info(f"XGBoost retrained: {len(data)} samples → {path}")
    return {"success": True, "samples": len(data)}