"""
Aggancia gli xG di Understat (data/processed/xg.csv) alle partite nel database.

I nomi squadra di Understat differiscono da quelli di football-data
("AC Milan" vs "Milan"): la mappatura viene imparata automaticamente
accoppiando le partite per (lega, data, risultato esatto) nei casi non
ambigui, poi applicata a tutto il dataset.

Aggiunge (se mancano) le colonne home_xg / away_xg alla tabella matches
e le popola. Idempotente.

Uso:
    python scripts/import_xg.py
"""

import argparse
import sys
from collections import Counter, defaultdict
from pathlib import Path

import pandas as pd
import pymysql

import config

PROJECT_ROOT = Path(__file__).resolve().parent.parent
XG_PATH = PROJECT_ROOT / "data" / "processed" / "xg.csv"
MAP_PATH = PROJECT_ROOT / "data" / "processed" / "team_mapping.csv"


def main() -> int:
    parser = argparse.ArgumentParser(description="Importa xG Understat nel DB")
    config.add_db_args(parser)
    args = parser.parse_args()

    xg = pd.read_csv(XG_PATH, parse_dates=["date"])
    xg["date"] = xg["date"].dt.normalize()

    conn = pymysql.connect(host=args.host, port=args.port, user=args.user,
                           password=args.password, database="tigertips")
    cur = conn.cursor()

    # Colonne xG (se mancano)
    cur.execute("""
        SELECT COUNT(*) FROM information_schema.columns
        WHERE table_schema = 'tigertips' AND table_name = 'matches'
          AND column_name = 'home_xg'
    """)
    if cur.fetchone()[0] == 0:
        cur.execute("ALTER TABLE matches "
                    "ADD COLUMN home_xg DECIMAL(6,3) NULL AFTER away_red, "
                    "ADD COLUMN away_xg DECIMAL(6,3) NULL AFTER home_xg")
        conn.commit()
        print("Colonne home_xg/away_xg aggiunte alla tabella matches")

    cur.execute("""
        SELECT m.id, l.code, m.match_date, h.name, a.name, m.home_goals, m.away_goals
        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 l.code IN ('E0', 'SP1', 'D1', 'I1', 'F1')
    """)
    db = pd.DataFrame(cur.fetchall(), columns=[
        "id", "league", "date", "home_db", "away_db", "hg", "ag"])
    db["date"] = pd.to_datetime(db["date"])

    # --- Fase 1: impara la mappatura nomi con accoppiamenti non ambigui ---
    key = ["league", "date", "hg", "ag"]
    db_unique = db.groupby(key).filter(lambda g: len(g) == 1)
    xg_unique = xg.groupby(key).filter(lambda g: len(g) == 1)
    paired = xg_unique.merge(db_unique, on=key)

    votes: dict[str, Counter] = defaultdict(Counter)
    for r in paired.itertuples(index=False):
        votes[r.home_us][r.home_db] += 1
        votes[r.away_us][r.away_db] += 1

    name_map = {}
    for us_name, counter in votes.items():
        best, n_best = counter.most_common(1)[0]
        total = sum(counter.values())
        if n_best / total < 0.9:
            print(f"  [ATTENZIONE] mappatura ambigua per '{us_name}': {dict(counter)}")
        name_map[us_name] = best

    pd.DataFrame(sorted(name_map.items()), columns=["understat", "football_data"]) \
        .to_csv(MAP_PATH, index=False)
    print(f"Mappatura nomi: {len(name_map)} squadre (salvata in {MAP_PATH.name})")

    unmapped = (set(xg["home_us"]) | set(xg["away_us"])) - set(name_map)
    if unmapped:
        print(f"  [ATTENZIONE] squadre Understat senza mappatura: {sorted(unmapped)}")

    # --- Fase 2: join completo usando la mappatura (data esatta, poi +/- 1 giorno) ---
    xg["home_db"] = xg["home_us"].map(name_map)
    xg["away_db"] = xg["away_us"].map(name_map)
    xg_ok = xg.dropna(subset=["home_db", "away_db"]).copy()

    matched_parts = []
    remaining = xg_ok.reset_index(drop=True).reset_index(names="xg_idx")
    db_renamed = db.rename(columns={"date": "join_date"})
    for shift in (0, 1, -1):
        if remaining.empty:
            break
        attempt = remaining.copy()
        attempt["join_date"] = attempt["date"] + pd.Timedelta(days=shift)
        merged = attempt.merge(db_renamed,
                               on=["league", "join_date", "home_db", "away_db"],
                               suffixes=("_us", "_db"))
        matched_parts.append(merged)
        remaining = remaining[~remaining["xg_idx"].isin(merged["xg_idx"])]

    matched = pd.concat(matched_parts, ignore_index=True)
    matched = matched.drop_duplicates(subset=["id"], keep="first")

    # Verifica coerenza risultato
    bad = matched[(matched["hg_us"] != matched["hg_db"]) |
                  (matched["ag_us"] != matched["ag_db"])]
    if len(bad):
        print(f"  [ATTENZIONE] {len(bad)} accoppiamenti con risultato diverso, scartati")
        matched = matched.drop(index=bad.index)

    # --- Aggiornamento DB ---
    rows = [(float(r.xg_h), float(r.xg_a), int(r.id))
            for r in matched.itertuples(index=False)]
    batch = 1000
    for start in range(0, len(rows), batch):
        cur.executemany("UPDATE matches SET home_xg = %s, away_xg = %s WHERE id = %s",
                        rows[start:start + batch])
        conn.commit()

    print(f"\nPartite Understat: {len(xg)} | agganciate al DB: {len(matched)} "
          f"({len(matched) / len(xg):.1%})")
    cur.execute("""
        SELECT l.code, COUNT(*), SUM(m.home_xg IS NOT NULL)
        FROM matches m JOIN leagues l ON l.id = m.league_id
        GROUP BY l.code ORDER BY l.code
    """)
    print("Copertura xG nel DB:")
    for code, tot, with_xg in cur.fetchall():
        print(f"  {code:<5} {int(with_xg or 0):>6} / {tot}")

    conn.close()
    return 0


if __name__ == "__main__":
    sys.exit(main())
