"""
Punto 2 della roadmap: backtest walk-forward del modello base Dixon-Coles,
misurato contro le probabilità implicite delle quote di chiusura.

Per ogni lega, il modello viene ri-addestrato periodicamente (default: ogni 7
giorni) usando SOLO le partite precedenti, e predice le partite del periodo
successivo. Nessun dato del futuro entra mai nel training (no look-ahead).

Metriche: log loss e Brier score multiclasse, confrontati con:
  - il mercato (probabilità implicite delle quote di chiusura, senza margine)
  - un baseline naive (frequenze storiche H/D/A della lega)

Uso:
    python scripts/backtest.py
    python scripts/backtest.py --eval-seasons 2023/24 2024/25 2025/26 --refit-days 14

Output: data/processed/backtest_base.csv (una riga per partita predetta)
"""

import argparse
import sys
import time
from pathlib import Path

import numpy as np
import pandas as pd
import pymysql

import config

import dixon_coles as dc

PROJECT_ROOT = Path(__file__).resolve().parent.parent
OUT_PATH = PROJECT_ROOT / "data" / "processed" / "backtest_base.csv"

XI = 0.0019            # decadimento temporale: peso = exp(-XI * giorni fa)
TRAIN_MAX_DAYS = 365 * 5  # oltre 5 anni il peso è < 0.03: inutile tenerle
RESULT_INDEX = {"H": 0, "D": 1, "A": 2}


def load_matches(args) -> pd.DataFrame:
    conn = pymysql.connect(host=args.host, port=args.port, user=args.user,
                           password=args.password, database="tigertips")
    cur = conn.cursor()
    cur.execute("""
        SELECT l.code, m.season, m.match_date, h.name, a.name,
               m.home_goals, m.away_goals, m.result,
               m.imp_prob_home, m.imp_prob_draw, m.imp_prob_away,
               m.home_xg, m.away_xg
        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
        ORDER BY m.match_date
    """)
    df = pd.DataFrame(cur.fetchall(), columns=[
        "league", "season", "date", "home", "away",
        "hg", "ag", "result", "mkt_h", "mkt_d", "mkt_a", "xg_h", "xg_a"])
    conn.close()
    df["date"] = pd.to_datetime(df["date"])
    for c in ("mkt_h", "mkt_d", "mkt_a", "xg_h", "xg_a"):
        df[c] = df[c].astype(float)

    # Pseudo-gol per il training: mix gol reali / xG (dove l'xG manca: solo gol)
    alpha = args.alpha
    df["hg_train"] = np.where(df["xg_h"].notna(),
                              (1 - alpha) * df["hg"] + alpha * df["xg_h"], df["hg"])
    df["ag_train"] = 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) -> pd.DataFrame:
    data = df[df["league"] == league].reset_index(drop=True)
    eval_mask = data["season"].isin(eval_seasons)
    if not eval_mask.any():
        return pd.DataFrame()

    eval_dates = sorted(data.loc[eval_mask, "date"].unique())
    window_start = eval_dates[0]
    last_date = eval_dates[-1]

    rows = []
    warm_start = None
    prev_teams: list[str] | None = None
    n_fits = 0
    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))]
        to_predict = data[eval_mask & (data["date"] >= window_start) &
                          (data["date"] < window_end)]
        if to_predict.empty:
            window_start = window_end
            continue

        teams = sorted(set(train["home"]) | set(train["away"]))
        team_idx = {t: i for i, t in enumerate(teams)}
        days_ago = (window_start - train["date"]).dt.days.to_numpy()
        weights = np.exp(-XI * days_ago)

        # warm start valido solo se l'insieme squadre non è cambiato
        x0 = warm_start if teams == prev_teams else None
        model = dc.fit(
            home_idx=train["home"].map(team_idx).to_numpy(),
            away_idx=train["away"].map(team_idx).to_numpy(),
            home_goals=train["hg_train"].to_numpy(),
            away_goals=train["ag_train"].to_numpy(),
            weights=weights, teams=teams, x0=x0,
        )
        warm_start, prev_teams = model.params_vector(), teams
        n_fits += 1

        for r in to_predict.itertuples(index=False):
            att_h, dfn_h = dc.team_params(model, r.home)
            att_a, dfn_a = dc.team_params(model, r.away)
            p_h, p_d, p_a = dc.predict(model, att_h, dfn_h, att_a, dfn_a)
            rows.append({
                "league": league, "season": r.season, "date": r.date,
                "home": r.home, "away": r.away, "result": r.result,
                "model_h": p_h, "model_d": p_d, "model_a": p_a,
                "mkt_h": r.mkt_h, "mkt_d": r.mkt_d, "mkt_a": r.mkt_a,
            })
        window_start = window_end

    print(f"  {league}: {len(rows)} partite predette, {n_fits} fit, "
          f"{time.time() - t0:.0f}s")
    return pd.DataFrame(rows)


def metrics(probs: np.ndarray, outcome_idx: np.ndarray) -> dict:
    """Log loss, Brier multiclasse e accuratezza."""
    p = np.clip(probs, 1e-12, 1)
    p = p / p.sum(axis=1, keepdims=True)
    n = len(outcome_idx)
    onehot = np.zeros_like(p)
    onehot[np.arange(n), outcome_idx] = 1
    return {
        "logloss": float(-np.mean(np.log(p[np.arange(n), outcome_idx]))),
        "brier": float(np.mean(np.sum((p - onehot) ** 2, axis=1))),
        "acc": float(np.mean(p.argmax(axis=1) == outcome_idx)),
    }


def report(bt: pd.DataFrame, train_freqs: dict[str, np.ndarray]) -> None:
    scored = bt.dropna(subset=["mkt_h", "mkt_d", "mkt_a"]).copy()
    outcome = scored["result"].map(RESULT_INDEX).to_numpy()
    model_p = scored[["model_h", "model_d", "model_a"]].to_numpy()
    mkt_p = scored[["mkt_h", "mkt_d", "mkt_a"]].to_numpy()
    naive_p = np.vstack([train_freqs[lg] for lg in scored["league"]])

    print(f"\n=== Risultati backtest ({len(scored)} partite con quote) ===")
    header = f"{'':<18}{'log loss':>10}{'Brier':>10}{'accur.':>9}"
    print(header)
    for name, p in [("Mercato (closing)", mkt_p),
                    ("Modello DC", model_p),
                    ("Naive (frequenze)", naive_p)]:
        m = metrics(p, outcome)
        print(f"{name:<18}{m['logloss']:>10.4f}{m['brier']:>10.4f}{m['acc']:>9.1%}")

    print("\nLog loss per lega (modello vs mercato):")
    for lg, grp in scored.groupby("league"):
        oi = grp["result"].map(RESULT_INDEX).to_numpy()
        m_model = metrics(grp[["model_h", "model_d", "model_a"]].to_numpy(), oi)
        m_mkt = metrics(grp[["mkt_h", "mkt_d", "mkt_a"]].to_numpy(), oi)
        gap = m_model["logloss"] - m_mkt["logloss"]
        print(f"  {lg:<5} modello {m_model['logloss']:.4f}   "
              f"mercato {m_mkt['logloss']:.4f}   gap {gap:+.4f}")


def main() -> int:
    parser = argparse.ArgumentParser(description="Backtest walk-forward Dixon-Coles")
    parser.add_argument("--eval-seasons", nargs="+", default=["2024/25", "2025/26"],
                        help="Stagioni da predire (default: 2024/25 2025/26)")
    parser.add_argument("--refit-days", type=int, default=7,
                        help="Ogni quanti giorni ri-addestrare (default: 7)")
    parser.add_argument("--leagues", nargs="+", default=None,
                        help="Codici lega da testare (default: tutte)")
    parser.add_argument("--alpha", type=float, default=0.0,
                        help="Peso xG negli pseudo-gol di training: "
                             "0 = solo gol reali (default), 1 = solo xG")
    parser.add_argument("--out", default=None,
                        help="File di output (default: backtest_base.csv)")
    config.add_db_args(parser)
    args = parser.parse_args()

    df = load_matches(args)
    leagues = args.leagues or sorted(df["league"].unique())
    out_path = Path(args.out) if args.out else OUT_PATH
    print(f"Backtest su {leagues}, stagioni {args.eval_seasons}, "
          f"refit ogni {args.refit_days} giorni, alpha xG = {args.alpha}")

    # Frequenze H/D/A per lega calcolate SOLO sulle stagioni di training
    train_only = df[~df["season"].isin(args.eval_seasons)]
    train_freqs = {}
    for lg, grp in train_only.groupby("league"):
        counts = grp["result"].map(RESULT_INDEX).value_counts().reindex([0, 1, 2]).fillna(0)
        train_freqs[lg] = (counts / counts.sum()).to_numpy()

    parts = [backtest_league(df, lg, args.eval_seasons, args.refit_days)
             for lg in leagues]
    bt = pd.concat([p for p in parts if not p.empty], ignore_index=True)

    out_path.parent.mkdir(parents=True, exist_ok=True)
    bt.to_csv(out_path, index=False)
    report(bt, train_freqs)
    print(f"\nPredizioni salvate in: {out_path}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
