"""
Scoring CLV dell'agente formazioni (e futuri agenti live).

Per ogni segnale in agent_signals con partita giocata:
  - snapshot: le probabilità di mercato più recenti PRIMA del segnale
    (odds_snapshots; fallback: quote pre-match in matches)
  - closing: le probabilità implicite di chiusura (imp_prob_* in matches)
  - direzione: sull'esito dove l'agente ha spostato di più la probabilità,
    la closing si è mossa nella stessa direzione dello spostamento?

KPI primario (dal design): CLV direction hit rate. >55% = l'agente vede cose
vere prima del mercato; <=52% dopo 150 segnali non nulli = kill criterion.
KPI secondario: log loss delle probabilità aggiustate vs snapshot vs closing.

Uso:
    python scripts/score_signals.py
"""

import argparse
import sys

import numpy as np
import pandas as pd
import pymysql

import config

MIN_DELTA = 0.005      # segnale "non nullo": aggiustamento lambda oltre questa soglia
FLAT_MOVE = 0.002      # movimento closing sotto questa soglia = piatto (escluso)
RESULT_INDEX = {"H": 0, "D": 1, "A": 2}


def logloss(p: np.ndarray, outcome: np.ndarray) -> float:
    p = np.clip(p, 1e-12, 1)
    p = p / p.sum(axis=1, keepdims=True)
    return float(-np.mean(np.log(p[np.arange(len(outcome)), outcome])))


def main() -> int:
    parser = argparse.ArgumentParser(description="Scoring CLV dei segnali agente")
    config.add_db_args(parser)
    args = parser.parse_args()

    conn = pymysql.connect(host=args.host, port=args.port, user=args.user,
                           password=args.password, database="tigertips")

    df = pd.read_sql("""
        SELECT s.id AS signal_id, s.agent, s.stage, s.checked_at,
               s.delta_lambda_home, s.delta_lambda_away,
               m.result,
               m.imp_prob_home AS close_h, m.imp_prob_draw AS close_d,
               m.imp_prob_away AS close_a,
               pm.prob_home AS model_h, pm.prob_draw AS model_d, pm.prob_away AS model_a,
               pa.prob_home AS adj_h, pa.prob_draw AS adj_d, pa.prob_away AS adj_a,
               os.prob_home AS snap_h, os.prob_draw AS snap_d, os.prob_away AS snap_a
        FROM agent_signals s
        JOIN matches m ON m.id = s.match_id AND m.result IS NOT NULL
        LEFT JOIN predictions pm ON pm.match_id = s.match_id
             AND pm.stage = s.stage AND pm.source = 'model'
        LEFT JOIN predictions pa ON pa.match_id = s.match_id
             AND pa.stage = s.stage AND pa.source = 'agent_adj'
        LEFT JOIN odds_snapshots os ON os.id = (
            SELECT os2.id FROM odds_snapshots os2
            WHERE os2.match_id = s.match_id AND os2.taken_at <= s.checked_at
            ORDER BY os2.taken_at DESC LIMIT 1)
    """, conn)
    conn.close()

    if df.empty:
        print("Nessun segnale con partita giocata: lo scoring parte con la stagione.")
        return 0

    for c in df.columns.drop(["agent", "stage", "checked_at", "result"]):
        df[c] = pd.to_numeric(df[c], errors="coerce")

    df["nonnull"] = (df["delta_lambda_home"].abs() +
                     df["delta_lambda_away"].abs()) > MIN_DELTA

    print(f"Segnali con esito: {len(df)} | non nulli: {df['nonnull'].sum()}\n")

    for (agent, stage), grp in df.groupby(["agent", "stage"]):
        print(f"=== {agent} / {stage} ({len(grp)} segnali, "
              f"{grp['nonnull'].sum()} non nulli) ===")

        # --- CLV direzionale sui segnali non nulli con snapshot e closing ---
        g = grp[grp["nonnull"]].dropna(
            subset=["snap_h", "adj_h", "close_h"]).copy()
        if len(g):
            snap = g[["snap_h", "snap_d", "snap_a"]].to_numpy()
            adj = g[["adj_h", "adj_d", "adj_a"]].to_numpy()
            close = g[["close_h", "close_d", "close_a"]].to_numpy()
            shift = adj - snap
            comp = np.abs(shift).argmax(axis=1)
            idx = np.arange(len(g))
            agent_dir = np.sign(shift[idx, comp])
            close_move = close[idx, comp] - snap[idx, comp]
            flat = np.abs(close_move) < FLAT_MOVE
            hits = np.sign(close_move[~flat]) == agent_dir[~flat]
            if hits.size:
                rate = hits.mean()
                verdict = ("POSITIVO: l'agente anticipa il mercato" if rate > 0.55
                           else "in linea col rumore" if rate > 0.48
                           else "NEGATIVO")
                print(f"  CLV direzionale: {rate:.1%} su {hits.size} segnali "
                      f"(piatti esclusi: {flat.sum()}) -> {verdict}")
            if hits.size >= 150 and hits.mean() <= 0.52:
                print("  *** KILL CRITERION RAGGIUNTO: valutare archiviazione ***")
        else:
            print("  CLV direzionale: dati insufficienti (serve snapshot+closing)")

        # --- Log loss comparativo su tutti i segnali con esito ---
        g2 = grp.dropna(subset=["adj_h", "close_h", "result"]).copy()
        g2 = g2[g2["result"].isin(RESULT_INDEX)]
        if len(g2) >= 10:
            y = g2["result"].map(RESULT_INDEX).to_numpy()
            print(f"  log loss ({len(g2)} partite): "
                  f"aggiustate {logloss(g2[['adj_h', 'adj_d', 'adj_a']].to_numpy(), y):.4f} | "
                  f"closing {logloss(g2[['close_h', 'close_d', 'close_a']].to_numpy(), y):.4f}")
        print()
    return 0


if __name__ == "__main__":
    sys.exit(main())
