"""
Backtest storico delle tre fasce dei pronostici del giorno.

Stessa meccanica walk-forward di backtest.py (refit settimanale, zero
look-ahead), ma per ogni partita: griglia punteggi DC -> riscalata alle
probabilità implicite del mercato -> selezione delle tre fasce -> esito reale.

Verifica la promessa di calibrazione: la fascia "sicuro" (p dichiarata
0.70-0.88) vince davvero in quella banda? Idem per le altre.

Uso:
    python scripts/backtest_picks.py
"""

import argparse
import sys
import time

import numpy as np
import pandas as pd
import pymysql

import config

import dixon_coles as dc
from score_tipsters import settle
from daily_picks import candidates_from_grid, pick_tiers

XI = 0.0019
ALPHA = 0.7
TRAIN_MAX_DAYS = 365 * 5


def load_matches(args) -> pd.DataFrame:
    conn = pymysql.connect(host=args.host, port=args.port, user=args.user,
                           password=args.password, database="tigertips")
    df = pd.read_sql("""
        SELECT l.code AS league, m.season, m.match_date AS date,
               h.name AS home, a.name AS away,
               m.home_goals AS hg, m.away_goals AS ag,
               m.imp_prob_home AS mkt_h, m.imp_prob_draw AS mkt_d,
               m.imp_prob_away AS mkt_a,
               m.home_xg AS xg_h, m.away_xg AS xg_a
        FROM matches m
        JOIN leagues l ON l.id = m.league_id
        JOIN teams h ON h.id = m.home_team_id
        JOIN teams a ON a.id = m.away_team_id
        WHERE m.result IS NOT NULL
        ORDER BY m.match_date
    """, conn)
    conn.close()
    df["date"] = pd.to_datetime(df["date"])
    for c in ("mkt_h", "mkt_d", "mkt_a", "xg_h", "xg_a"):
        df[c] = pd.to_numeric(df[c], errors="coerce")
    df["hg_t"] = np.where(df["xg_h"].notna(),
                          (1 - ALPHA) * df["hg"] + ALPHA * df["xg_h"], df["hg"])
    df["ag_t"] = np.where(df["xg_a"].notna(),
                          (1 - ALPHA) * df["ag"] + ALPHA * df["xg_a"], df["ag"])
    return df


def backtest_league(df: pd.DataFrame, league: str, eval_seasons: list[str],
                    refit_days: int) -> list[dict]:
    data = df[df["league"] == league].reset_index(drop=True)
    eval_mask = data["season"].isin(eval_seasons)
    if not eval_mask.any():
        return []
    eval_dates = sorted(data.loc[eval_mask, "date"].unique())
    window_start, last_date = eval_dates[0], eval_dates[-1]

    rows = []
    warm, prev_teams = None, None
    t0 = time.time()
    while window_start <= last_date:
        window_end = window_start + pd.Timedelta(days=refit_days)
        train = data[(data["date"] < window_start) &
                     (data["date"] >= window_start - pd.Timedelta(days=TRAIN_MAX_DAYS))]
        todo = data[eval_mask & (data["date"] >= window_start) &
                    (data["date"] < window_end)]
        if todo.empty:
            window_start = window_end
            continue
        teams = sorted(set(train["home"]) | set(train["away"]))
        idx = {t: i for i, t in enumerate(teams)}
        days_ago = (window_start - train["date"]).dt.days.to_numpy()
        model = dc.fit(train["home"].map(idx).to_numpy(),
                       train["away"].map(idx).to_numpy(),
                       train["hg_t"].to_numpy(), train["ag_t"].to_numpy(),
                       np.exp(-XI * days_ago), teams,
                       x0=warm if teams == prev_teams else None)
        warm, prev_teams = model.params_vector(), teams

        for r in todo.itertuples(index=False):
            att_h, dfn_h = dc.team_params(model, r.home)
            att_a, dfn_a = dc.team_params(model, r.away)
            grid = dc.score_matrix(model, att_h, dfn_h, att_a, dfn_a)
            if pd.notna(r.mkt_h):
                grid = dc.rescale_to_1x2(grid, (r.mkt_h, r.mkt_d, r.mkt_a))
            for tier, cand in pick_tiers(candidates_from_grid(grid)).items():
                outcome = settle({"market": cand["market"],
                                  "selection": cand["selection"],
                                  "line": cand["line"], "hg": r.hg, "ag": r.ag})
                rows.append({"league": league, "tier": tier, "p": cand["p"],
                             "outcome": outcome})
        window_start = window_end
    print(f"  {league}: {time.time() - t0:.0f}s")
    return rows


def main() -> int:
    parser = argparse.ArgumentParser(description="Backtest delle fasce pronostici")
    parser.add_argument("--eval-seasons", nargs="+", default=["2024/25", "2025/26"])
    parser.add_argument("--refit-days", type=int, default=7)
    config.add_db_args(parser)
    args = parser.parse_args()

    df = load_matches(args)
    leagues = sorted(df["league"].unique())
    print(f"Backtest fasce su {leagues}, stagioni {args.eval_seasons}")

    rows = []
    for lg in leagues:
        rows.extend(backtest_league(df, lg, args.eval_seasons, args.refit_days))
    bt = pd.DataFrame(rows)
    bt = bt[bt["outcome"].isin(["win", "loss"])]

    print(f"\n=== Verifica fasce ({len(bt)} pick regolate) ===")
    print(f"{'fascia':<14}{'pick':>7}{'p dichiarata':>14}{'vinte davvero':>15}")
    for tier in ("sicuro", "equilibrato", "azzardo"):
        g = bt[bt["tier"] == tier]
        if len(g):
            print(f"{tier:<14}{len(g):>7}{g['p'].mean():>14.1%}"
                  f"{(g['outcome'] == 'win').mean():>15.1%}")

    print("\nPer lega (fascia sicuro):")
    g = bt[bt["tier"] == "sicuro"]
    for lg, grp in g.groupby("league"):
        print(f"  {lg:<5} {(grp['outcome'] == 'win').mean():.1%} su {len(grp)}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
