import math
import random
from dataclasses import dataclass
from typing import Any, List, Mapping, Optional, Sequence, Tuple

from core import clamp, parse_mu_disp_mode, require, require_unit_interval_open_closed
from observables import EdgeCounts, mu_dispersion, mu_hat_by_comm_and_dispersion_from_counts, mu_risk_by_comm_and_dispersion_from_counts


def _dirichlet_sample(alpha: List[float], rng: random.Random) -> List[float]:
    xs: List[float] = []
    s = 0.0
    for a in alpha:
        aa = float(a)
        if aa <= 0.0:
            aa = 1e-12
        x = rng.gammavariate(aa, 1.0)
        xs.append(x)
        s += x
    if s <= 0.0:
        k = len(alpha)
        return [1.0 / k for _ in range(k)]
    return [x / s for x in xs]


def _dirichlet_mean(alpha: List[float]) -> List[float]:
    s = float(sum(alpha))
    if s <= 0.0:
        k = len(alpha)
        return [1.0 / k for _ in range(k)]
    return [float(a) / s for a in alpha]


def _mu_eta_from_p(p_cc: float, p_cd: float, p_dd: float) -> Tuple[float, float, float]:
    mu = p_cc + 0.5 * p_cd
    eta = p_cd
    h = 0.0
    for x in (p_cc, p_cd, p_dd):
        if x > 0.0:
            h -= x * math.log(x)
    h /= math.log(3.0)
    return float(mu), float(eta), float(h)


@dataclass
class DetectorState:
    alpha: List[float]
    v_mu: float
    v_eta: float


@dataclass
class DetectorOutput:
    b: float
    risk: float
    b_mu: float
    b_eta: float
    b_v: float
    b_disp: float
    mu_hat: float
    eta_hat: float
    h_hat: float
    mu_disp_hat: float
    v_mu_hat: float
    v_eta_hat: float
    v_hat: float


class DirichletVolatilityDetector:
    def __init__(self, cfg: Mapping[str, Any], rng: random.Random, *, risk_rng: Optional[random.Random] = None):
        a0 = require(cfg, "alpha0", "detector")
        if not isinstance(a0, list) or len(a0) != 3:
            raise ValueError("detector.alpha0 must be a list of length 3")
        alpha0 = [float(x) for x in a0]
        if any(x <= 0.0 for x in alpha0):
            raise ValueError("detector.alpha0 must be positive")

        self.rho = require_unit_interval_open_closed(require(cfg, "rho", "detector"), "detector.rho")
        self.k_mc = int(require(cfg, "k_mc", "detector"))
        self._risk_k_mc_override = bool("risk_k_mc" in cfg)
        self.risk_k_mc = int(cfg.get("risk_k_mc", self.k_mc))
        self.alpha_v = require_unit_interval_open_closed(require(cfg, "alpha_v", "detector"), "detector.alpha_v")

        if self.k_mc <= 0:
            raise ValueError("detector.k_mc must be > 0")
        if self.risk_k_mc <= 0:
            raise ValueError("detector.risk_k_mc must be > 0")

        self._rng = rng
        self._risk_rng = risk_rng if risk_rng is not None else rng
        self.state = DetectorState(alpha=list(alpha0), v_mu=0.0, v_eta=0.0)
        self._prev_alpha = list(alpha0)

        a0_meso = cfg.get("meso_alpha0", a0)
        if not isinstance(a0_meso, list) or len(a0_meso) != 3:
            raise ValueError("detector.meso_alpha0 must be a list of length 3")
        meso_alpha0 = [float(x) for x in a0_meso]
        if any(x <= 0.0 for x in meso_alpha0):
            raise ValueError("detector.meso_alpha0 must be positive")
        self.meso_alpha0 = list(meso_alpha0)

    def _update_alpha(self, obs: EdgeCounts) -> None:
        a = self.state.alpha
        self._prev_alpha = list(a)
        self.state.alpha = [
            self.rho * float(a[0]) + float(obs.cc),
            self.rho * float(a[1]) + float(obs.cd),
            self.rho * float(a[2]) + float(obs.dd),
        ]

    def step(
        self,
        obs: EdgeCounts,
        ethics_cfg: Mapping[str, Any],
        *,
        comm_obs: Optional[Sequence[EdgeCounts]] = None,
        comm_min_total_edges: float = 1.0,
    ) -> DetectorOutput:
        mu_min = float(require(ethics_cfg, "mu_min", "ethics"))
        eta_max = float(require(ethics_cfg, "eta_max", "ethics"))
        v_max = float(require(ethics_cfg, "v_max", "ethics"))
        mu_disp_max = ethics_cfg.get("mu_disp_max")
        use_disp = bool(mu_disp_max is not None and comm_obs is not None)
        mu_disp_max_f = float(mu_disp_max) if mu_disp_max is not None else None

        alpha_before = list(self.state.alpha)
        v_mu_before = float(self.state.v_mu)
        v_eta_before = float(self.state.v_eta)

        self._update_alpha(obs)
        a_cur = self.state.alpha
        a_prev = self._prev_alpha

        p_hat = _dirichlet_mean(a_cur)
        mu_hat, eta_hat, h_hat = _mu_eta_from_p(p_hat[0], p_hat[1], p_hat[2])

        p_hat_prev = _dirichlet_mean(alpha_before)
        mu_hat_prev, eta_hat_prev, _ = _mu_eta_from_p(p_hat_prev[0], p_hat_prev[1], p_hat_prev[2])

        dmu_hat = abs(mu_hat - mu_hat_prev)
        deta_hat = abs(eta_hat - eta_hat_prev)

        v_mu_new = (1.0 - self.alpha_v) * v_mu_before + self.alpha_v * dmu_hat
        v_eta_new = (1.0 - self.alpha_v) * v_eta_before + self.alpha_v * deta_hat
        v_new = max(v_mu_new, v_eta_new)

        b_all = 0
        b_mu = 0
        b_eta = 0
        b_v = 0
        b_disp = 0

        mu_disp_hat = 0.0
        mu_disp_mode = "mu_hat"
        if mu_disp_max is not None:
            mu_disp_mode = parse_mu_disp_mode(ethics_cfg)
        if use_disp:
            if mu_disp_mode == "risk":
                risk_k_mc = int(self.risk_k_mc)
                if not bool(self._risk_k_mc_override):
                    risk_k_mc = int(ethics_cfg.get("risk_k_mc", risk_k_mc))
                if risk_k_mc <= 0:
                    raise ValueError("risk_k_mc must be > 0 (set detector.risk_k_mc or ethics.risk_k_mc)")
                _, mu_disp_hat = mu_risk_by_comm_and_dispersion_from_counts(
                    comm_obs,
                    meso_alpha0=self.meso_alpha0,
                    mu_min=float(mu_min),
                    min_total_edges=float(comm_min_total_edges),
                    k_mc=int(risk_k_mc),
                    rng=self._risk_rng,
                )
            else:
                _, mu_disp_hat = mu_hat_by_comm_and_dispersion_from_counts(
                    comm_obs,
                    meso_alpha0=self.meso_alpha0,
                    min_total_edges=float(comm_min_total_edges),
                )

        for _ in range(self.k_mc):
            p1 = _dirichlet_sample(a_cur, self._rng)
            p0 = _dirichlet_sample(a_prev, self._rng)

            mu1, eta1, _ = _mu_eta_from_p(p1[0], p1[1], p1[2])
            mu0, eta0, _ = _mu_eta_from_p(p0[0], p0[1], p0[2])

            ok_mu = mu1 >= mu_min
            ok_eta = eta1 <= eta_max

            dmu = abs(mu1 - mu0)
            deta = abs(eta1 - eta0)
            v_mu_s = (1.0 - self.alpha_v) * v_mu_before + self.alpha_v * dmu
            v_eta_s = (1.0 - self.alpha_v) * v_eta_before + self.alpha_v * deta
            v_s = v_mu_s if v_mu_s >= v_eta_s else v_eta_s
            ok_v = v_s <= v_max

            ok_disp = True
            if use_disp and mu_disp_mode == "risk":
                ok_disp = bool(float(mu_disp_hat) <= float(mu_disp_max_f))
            elif use_disp:
                mus: List[float] = []
                for c in comm_obs:
                    if float(c.total) < float(comm_min_total_edges):
                        continue
                    a_comm = [
                        float(self.meso_alpha0[0]) + float(c.cc),
                        float(self.meso_alpha0[1]) + float(c.cd),
                        float(self.meso_alpha0[2]) + float(c.dd),
                    ]
                    p_comm = _dirichlet_sample(a_comm, self._rng)
                    mu_c, _, _ = _mu_eta_from_p(p_comm[0], p_comm[1], p_comm[2])
                    mus.append(float(mu_c))
                mu_disp_s = float(mu_dispersion(mus))
                ok_disp = bool(mu_disp_s <= float(mu_disp_max_f))

            if ok_mu:
                b_mu += 1
            if ok_eta:
                b_eta += 1
            if ok_v:
                b_v += 1
            if ok_disp:
                b_disp += 1
            if ok_mu and ok_eta and ok_v and ok_disp:
                b_all += 1

        b = b_all / float(self.k_mc)
        b_mu_f = b_mu / float(self.k_mc)
        b_eta_f = b_eta / float(self.k_mc)
        b_v_f = b_v / float(self.k_mc)
        b_disp_f = b_disp / float(self.k_mc)

        self.state.v_mu = float(v_mu_new)
        self.state.v_eta = float(v_eta_new)

        risk = 1.0 - b
        risk = clamp(risk, 0.0, 1.0)

        return DetectorOutput(
            b=float(b),
            risk=float(risk),
            b_mu=float(b_mu_f),
            b_eta=float(b_eta_f),
            b_v=float(b_v_f),
            b_disp=float(b_disp_f if use_disp else 1.0),
            mu_hat=float(mu_hat),
            eta_hat=float(eta_hat),
            h_hat=float(h_hat),
            mu_disp_hat=float(mu_disp_hat),
            v_mu_hat=float(self.state.v_mu),
            v_eta_hat=float(self.state.v_eta),
            v_hat=float(v_new),
        )
