#!/usr/bin/env python3
"""過去試合をウォークフォワードでざっくりバックテストする。"""
from __future__ import annotations

import argparse
import os
import sys
from collections import defaultdict

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from ai.simple_score import predict_matchup, reset_season_forms, update_form  # noqa: E402
from collectors.db_util import connect  # noqa: E402


def load_games(season_from: int, season_to: int, game_type: str = "Reg Season"):
    sql = """
        SELECT
          g.game_id, g.season, g.game_date,
          g.home_team_id, g.away_team_id,
          g.home_score, g.away_score,
          th.name AS home_team, ta.name AS away_team
        FROM games g
        JOIN teams th ON th.id = g.home_team_id
        JOIN teams ta ON ta.id = g.away_team_id
        WHERE g.season BETWEEN %s AND %s
          AND g.game_type = %s
        ORDER BY g.season, g.game_date, g.game_id
    """
    with connect() as conn:
        with conn.cursor() as cur:
            cur.execute(sql, (season_from, season_to, game_type))
            return cur.fetchall()


def run_backtest(season_from: int, season_to: int) -> dict:
    games = load_games(season_from, season_to)
    if not games:
        return {"error": "no games", "total": 0}

    forms: dict[int, object] = {}
    current_season: int | None = None

    total = 0
    correct = 0
    brier_sum = 0.0
    by_confidence: dict[str, list[int]] = defaultdict(list)

    for g in games:
        season = int(g["season"])
        if current_season != season:
            team_ids = set()
            for row in games:
                if int(row["season"]) == season:
                    team_ids.add(int(row["home_team_id"]))
                    team_ids.add(int(row["away_team_id"]))
            forms = reset_season_forms(team_ids)
            current_season = season

        home_id = int(g["home_team_id"])
        away_id = int(g["away_team_id"])
        home_form = forms[home_id]
        away_form = forms[away_id]

        pred = predict_matchup(home_form, away_form)
        p_home = pred["home_win_prob"]
        home_won = int(g["home_score"]) > int(g["away_score"])
        tie = int(g["home_score"]) == int(g["away_score"])

        if not tie:
            total += 1
            predicted_home = p_home >= 0.5
            if predicted_home == home_won:
                correct += 1
            actual = 1.0 if home_won else 0.0
            brier_sum += (p_home - actual) ** 2

            margin = abs(p_home - 0.5)
            if margin >= 0.15:
                bucket = "high"
            elif margin >= 0.08:
                bucket = "mid"
            else:
                bucket = "low"
            by_confidence[bucket].append(1 if predicted_home == home_won else 0)

        update_form(home_form, int(g["home_score"]), int(g["away_score"]))
        update_form(away_form, int(g["away_score"]), int(g["home_score"]))

    def bucket_acc(bucket: str) -> str:
        vals = by_confidence.get(bucket, [])
        if not vals:
            return "n/a"
        return f"{sum(vals) / len(vals) * 100:.1f}% ({len(vals)})"

    return {
        "season_from": season_from,
        "season_to": season_to,
        "total_games": total,
        "accuracy": round(correct / total * 100, 2) if total else 0,
        "brier_score": round(brier_sum / total, 4) if total else None,
        "accuracy_high_conf": bucket_acc("high"),
        "accuracy_mid_conf": bucket_acc("mid"),
        "accuracy_low_conf": bucket_acc("low"),
    }


def main() -> None:
    p = argparse.ArgumentParser(description="ざっくり勝敗バックテスト")
    p.add_argument("--from-season", type=int, default=2018)
    p.add_argument("--to-season", type=int, default=2024)
    args = p.parse_args()

    result = run_backtest(args.from_season, args.to_season)
    if result.get("error"):
        print(f"ERROR: {result['error']}")
        sys.exit(1)

    print(f"期間: {result['season_from']}〜{result['season_to']}")
    print(f"試合数: {result['total_games']}")
    print(f"勝敗的中率: {result['accuracy']}%")
    print(f"Brier score: {result['brier_score']} (低いほど良い)")
    print(f"  高信頼(差15%超): {result['accuracy_high_conf']}")
    print(f"  中信頼(差8〜15%): {result['accuracy_mid_conf']}")
    print(f"  低信頼(差8%未満): {result['accuracy_low_conf']}")


if __name__ == "__main__":
    main()
