"""
Modello Dixon-Coles (1997): Poisson bivariato con correzione per i punteggi
bassi (rho) e ponderazione temporale esponenziale delle partite di training.

Parametri stimati per lega:
    att[i]   - forza offensiva della squadra i
    dfn[i]   - forza difensiva della squadra i
    mu       - livello medio dei gol
    home_adv - vantaggio casa (sui gol attesi della squadra di casa)
    rho      - correzione di dipendenza per i punteggi 0-0, 1-0, 0-1, 1-1
"""

from dataclasses import dataclass

import numpy as np
from scipy.optimize import minimize
from scipy.special import gammaln

MAX_GOALS = 10          # griglia punteggi 0..10 per il calcolo delle probabilità
RHO_BOUNDS = (-0.15, 0.15)
L2_REG = 0.5            # piccola penalità ridge su att/dfn per stabilità


@dataclass
class DixonColesFit:
    teams: list[str]            # ordine dei parametri
    att: np.ndarray
    dfn: np.ndarray
    mu: float
    home_adv: float
    rho: float

    def params_vector(self) -> np.ndarray:
        return np.concatenate([self.att, self.dfn, [self.mu, self.home_adv, self.rho]])


def _unpack(params: np.ndarray, n: int):
    att = params[:n]
    dfn = params[n:2 * n]
    mu, home_adv, rho = params[2 * n], params[2 * n + 1], params[2 * n + 2]
    return att, dfn, mu, home_adv, rho


def _tau_log(hg, ag, lam_h, lam_a, rho):
    """log della correzione Dixon-Coles, vettorizzato sulle partite."""
    tau = np.ones_like(lam_h)
    m00 = (hg == 0) & (ag == 0)
    m01 = (hg == 0) & (ag == 1)
    m10 = (hg == 1) & (ag == 0)
    m11 = (hg == 1) & (ag == 1)
    tau[m00] = 1 - lam_h[m00] * lam_a[m00] * rho
    tau[m01] = 1 + lam_h[m01] * rho
    tau[m10] = 1 + lam_a[m10] * rho
    tau[m11] = 1 - rho
    return np.log(np.clip(tau, 1e-10, None))


def fit(home_idx: np.ndarray, away_idx: np.ndarray,
        home_goals: np.ndarray, away_goals: np.ndarray,
        weights: np.ndarray, teams: list[str],
        x0: np.ndarray | None = None) -> DixonColesFit:
    """Stima i parametri massimizzando la log-likelihood pesata."""
    n = len(teams)
    hg = home_goals.astype(float)
    ag = away_goals.astype(float)
    log_fact = gammaln(hg + 1) + gammaln(ag + 1)

    def nll(params):
        att, dfn, mu, home_adv, rho = _unpack(params, n)
        # centratura per identificabilità (la media di att e dfn è ridondante con mu)
        att = att - att.mean()
        dfn = dfn - dfn.mean()
        lam_h = np.exp(mu + home_adv + att[home_idx] - dfn[away_idx])
        lam_a = np.exp(mu + att[away_idx] - dfn[home_idx])
        ll = (-lam_h + hg * np.log(lam_h) - lam_a + ag * np.log(lam_a) - log_fact
              + _tau_log(hg, ag, lam_h, lam_a, rho))
        penalty = L2_REG * (np.sum(att ** 2) + np.sum(dfn ** 2))
        return -np.sum(weights * ll) + penalty

    if x0 is None:
        x0 = np.zeros(2 * n + 3)
        x0[2 * n] = np.log(max(hg.mean(), 0.5))   # mu iniziale dal dato
        x0[2 * n + 1] = 0.25                       # vantaggio casa tipico

    bounds = [(None, None)] * (2 * n) + [(None, None), (None, None), RHO_BOUNDS]
    res = minimize(nll, x0, method="L-BFGS-B", bounds=bounds,
                   options={"maxiter": 500})

    att, dfn, mu, home_adv, rho = _unpack(res.x, n)
    return DixonColesFit(teams=teams, att=att - att.mean(), dfn=dfn - dfn.mean(),
                         mu=mu, home_adv=home_adv, rho=rho)


def score_matrix(model: DixonColesFit, att_h: float, dfn_h: float,
                 att_a: float, dfn_a: float) -> np.ndarray:
    """Griglia P(gol_casa=i, gol_ospite=j) per i,j in 0..MAX_GOALS."""
    lam_h = np.exp(model.mu + model.home_adv + att_h - dfn_a)
    lam_a = np.exp(model.mu + att_a - dfn_h)

    goals = np.arange(MAX_GOALS + 1)
    pois_h = np.exp(-lam_h) * lam_h ** goals / np.exp(gammaln(goals + 1))
    pois_a = np.exp(-lam_a) * lam_a ** goals / np.exp(gammaln(goals + 1))
    grid = np.outer(pois_h, pois_a)

    rho = model.rho
    grid[0, 0] *= max(1 - lam_h * lam_a * rho, 1e-10)
    grid[0, 1] *= max(1 + lam_h * rho, 1e-10)
    grid[1, 0] *= max(1 + lam_a * rho, 1e-10)
    grid[1, 1] *= max(1 - rho, 1e-10)
    return grid / grid.sum()


def rescale_to_1x2(grid: np.ndarray, target: tuple[float, float, float]) -> np.ndarray:
    """
    Riscala le tre regioni della griglia (vittoria casa / pari / vittoria
    ospite) perché le loro somme coincidano con le probabilità target (es.
    quelle implicite del mercato). I mercati derivati (over, multigol...)
    diventano così coerenti con le quote 1X2.
    """
    grid = grid.copy()
    home_mask = np.tril(np.ones_like(grid, dtype=bool), -1)
    draw_mask = np.eye(grid.shape[0], dtype=bool)
    away_mask = np.triu(np.ones_like(grid, dtype=bool), 1)
    for mask, t in zip((home_mask, draw_mask, away_mask), target):
        s = grid[mask].sum()
        if s > 0:
            grid[mask] *= t / s
    return grid / grid.sum()


def predict(model: DixonColesFit, att_h: float, dfn_h: float,
            att_a: float, dfn_a: float) -> tuple[float, float, float]:
    """Probabilità (H, D, A) per una partita, dati i parametri delle due squadre."""
    grid = score_matrix(model, att_h, dfn_h, att_a, dfn_a)
    p_home = np.tril(grid, -1).sum()   # gol casa > gol ospite
    p_draw = np.trace(grid)
    p_away = np.triu(grid, 1).sum()
    return float(p_home), float(p_draw), float(p_away)


def team_params(model: DixonColesFit, team: str) -> tuple[float, float]:
    """Parametri (att, dfn) di una squadra; prior da neopromossa se sconosciuta."""
    if team in model.teams:
        i = model.teams.index(team)
        return float(model.att[i]), float(model.dfn[i])
    # Squadra mai vista nel training (es. neopromossa): assumiamo una squadra
    # debole, al 20° percentile di attacco e difesa della lega
    return float(np.quantile(model.att, 0.20)), float(np.quantile(model.dfn, 0.20))
