"""
Value at Risk (VaR) and Expected Shortfall (ES) Risk Measures.

Implements coherent risk measures for cryptocurrency portfolios:
- VaR: Value at Risk (parametric and historical)
- CVaR/ES: Conditional Value at Risk / Expected Shortfall

Mathematical Foundation:
- VaR_α(X) = -F_X^{-1}(α) = -inf{x : P(X ≤ x) ≥ α}
- ES_α(X) = E[X | X ≤ -VaR_α] = (1/α)∫₀^α VaR_u du

Parametric VaR (Student-t):
VaR_α = -μ - σ × √((ν-2)/ν) × t_ν^{-1}(α)

Artzner Coherence Axioms (ES satisfies all, VaR fails subadditivity):
1. Translation invariance: ρ(X + c) = ρ(X) - c
2. Subadditivity: ρ(X + Y) ≤ ρ(X) + ρ(Y)
3. Positive homogeneity: ρ(λX) = λρ(X)
4. Monotonicity: X ≤ Y a.s. ⟹ ρ(X) ≥ ρ(Y)

Crypto-specific values:
- 99% daily VaR: -13% to -18%
- 99% ES: -20% to -27%
- Student-t df: 3-6 (heavy tails)
"""
import math
from collections import deque
from dataclasses import dataclass
from typing import TYPE_CHECKING

from models import Tick
from ..base import BaseModel, SignalType, SignalCallback

if TYPE_CHECKING:
    from ..volatility.garch import GARCH11


@dataclass(slots=True)
class RiskMetrics:
    """Container for risk measure outputs."""
    var_95: float           # 95% VaR
    var_99: float           # 99% VaR
    es_95: float            # 95% Expected Shortfall
    es_99: float            # 99% Expected Shortfall
    current_return: float   # Latest return
    mean_return: float      # Mean return
    std_return: float       # Volatility


class RealTimeVaR(BaseModel):
    """
    Real-time VaR and Expected Shortfall calculator.

    Implements three methods:
    1. Parametric (Student-t): Fast, uses GARCH volatility if available
    2. Historical: Non-parametric, requires return history
    3. Cornish-Fisher: Adjusts for skewness and kurtosis

    Can integrate with GARCH model for volatility input.

    Emits RISK_BREACH signals when returns exceed VaR threshold.

    Attributes:
        symbol: Asset symbol to track
        garch: Optional GARCH model for volatility
        _returns: Deque of historical returns
        _nu: Student-t degrees of freedom
    """

    __slots__ = (
        'symbol', 'garch', 'confidence_levels',
        '_returns', '_nu', '_prev_price',
        '_mean_ema', '_var_ema', '_ema_decay',
        '_skew_ema', '_kurt_ema',
        '_var_breach_count', '_es_breach_count',
    )

    # Student-t quantiles for common confidence levels
    # Pre-computed for degrees of freedom 3-10
    T_QUANTILES = {
        # (nu, alpha): quantile
        (3, 0.95): 2.353,
        (3, 0.99): 4.541,
        (4, 0.95): 2.132,
        (4, 0.99): 3.747,
        (5, 0.95): 2.015,
        (5, 0.99): 3.365,
        (6, 0.95): 1.943,
        (6, 0.99): 3.143,
    }

    def __init__(
        self,
        symbol: str,
        confidence_levels: list[float] | None = None,
        history_size: int = 1000,
        nu: float = 5.0,
        garch_model: "GARCH11 | None" = None,
        warmup_ticks: int = 500,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize VaR/ES calculator.

        Args:
            symbol: Asset symbol to track
            confidence_levels: VaR confidence levels (default [0.95, 0.99])
            history_size: Number of returns to store for historical VaR
            nu: Student-t degrees of freedom (crypto: 3-6)
            garch_model: Optional GARCH model for volatility
            warmup_ticks: Ticks before stable estimates
            on_signal: Async callback
        """
        super().__init__(f"var_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol
        self.confidence_levels = confidence_levels or [0.95, 0.99]
        self.garch = garch_model

        # Validate nu
        if nu <= 2:
            raise ValueError("Student-t df must be > 2 for finite variance")
        self._nu = nu

        # Return history for historical VaR
        self._returns: deque = deque(maxlen=history_size)
        self._prev_price = 0.0

        # EMA statistics for parametric VaR
        self._ema_decay = 2.0 / (min(history_size, 200) + 1)
        self._mean_ema = 0.0
        self._var_ema = 0.0
        self._skew_ema = 0.0  # For Cornish-Fisher adjustment
        self._kurt_ema = 3.0  # Normal kurtosis

        # Breach tracking
        self._var_breach_count = 0
        self._es_breach_count = 0

    async def on_tick(self, tick: Tick) -> None:
        """Update risk measures with new tick."""
        if tick.symbol != self.symbol:
            return

        price = float(tick.mid_price())

        if self._prev_price <= 0:
            self._prev_price = price
            return

        # Log return
        ret = math.log(price / self._prev_price)
        self._returns.append(ret)
        self._prev_price = price

        # Update EMA statistics
        self._update_stats(ret)
        self._check_warmup()

        # Check for VaR breaches
        if self._is_ready:
            await self._check_breaches(ret)

    def _update_stats(self, ret: float) -> None:
        """Update EMA-based statistics."""
        # Mean
        delta = ret - self._mean_ema
        self._mean_ema += self._ema_decay * delta

        # Variance
        self._var_ema = (
            (1 - self._ema_decay) * self._var_ema +
            self._ema_decay * delta * delta
        )

        # Standardized return for higher moments
        std = math.sqrt(self._var_ema) if self._var_ema > 1e-16 else 1e-8
        z = delta / std

        # Skewness (EMA of z³)
        self._skew_ema = (
            (1 - self._ema_decay) * self._skew_ema +
            self._ema_decay * z * z * z
        )

        # Excess kurtosis (EMA of z⁴ - 3)
        self._kurt_ema = (
            (1 - self._ema_decay) * self._kurt_ema +
            self._ema_decay * (z * z * z * z)
        )

    async def _check_breaches(self, ret: float) -> None:
        """Check if current return breaches VaR threshold."""
        var_99 = self.get_parametric_var(0.99)

        # VaR breach: return more negative than VaR
        if ret < var_99:
            self._var_breach_count += 1
            await self._emit_signal(
                signal_type=SignalType.RISK_BREACH,
                symbol=self.symbol,
                confidence=min(abs(ret / var_99) - 1, 1.0) if var_99 != 0 else 0.5,
                metadata={
                    "breach_type": "var_99",
                    "return": ret,
                    "var_99": var_99,
                    "severity": ret / var_99 if var_99 != 0 else 0,
                },
            )

    def _get_sigma(self) -> float:
        """Get volatility, using GARCH if available."""
        if self.garch and self.garch.is_ready():
            return self.garch.get_sigma()
        return math.sqrt(self._var_ema) if self._var_ema > 0 else 1e-8

    def _t_quantile(self, alpha: float) -> float:
        """
        Get Student-t quantile for given alpha.

        Uses pre-computed values or approximation.
        """
        nu_int = int(round(self._nu))
        nu_int = max(3, min(6, nu_int))  # Clamp to available values

        key = (nu_int, alpha)
        if key in self.T_QUANTILES:
            return self.T_QUANTILES[key]

        # Approximation for other values using normal quantile adjustment
        # This is a rough approximation; for production use scipy.stats.t
        z_alpha = self._normal_quantile(alpha)
        # Adjust for heavier tails
        adjustment = 1 + (1 / (4 * self._nu))
        return z_alpha * adjustment

    def _normal_quantile(self, p: float) -> float:
        """Approximate normal quantile using Beasley-Springer-Moro."""
        if p <= 0 or p >= 1:
            return 0.0

        # Coefficients for rational approximation
        a = [
            -3.969683028665376e+01,
            2.209460984245205e+02,
            -2.759285104469687e+02,
            1.383577518672690e+02,
            -3.066479806614716e+01,
            2.506628277459239e+00,
        ]
        b = [
            -5.447609879822406e+01,
            1.615858368580409e+02,
            -1.556989798598866e+02,
            6.680131188771972e+01,
            -1.328068155288572e+01,
        ]
        c = [
            -7.784894002430293e-03,
            -3.223964580411365e-01,
            -2.400758277161838e+00,
            -2.549732539343734e+00,
            4.374664141464968e+00,
            2.938163982698783e+00,
        ]
        d = [
            7.784695709041462e-03,
            3.224671290700398e-01,
            2.445134137142996e+00,
            3.754408661907416e+00,
        ]

        p_low = 0.02425
        p_high = 1 - p_low

        if p < p_low:
            q = math.sqrt(-2 * math.log(p))
            return (((((c[0]*q + c[1])*q + c[2])*q + c[3])*q + c[4])*q + c[5]) / \
                   ((((d[0]*q + d[1])*q + d[2])*q + d[3])*q + 1)
        elif p <= p_high:
            q = p - 0.5
            r = q * q
            return (((((a[0]*r + a[1])*r + a[2])*r + a[3])*r + a[4])*r + a[5]) * q / \
                   (((((b[0]*r + b[1])*r + b[2])*r + b[3])*r + b[4])*r + 1)
        else:
            q = math.sqrt(-2 * math.log(1 - p))
            return -(((((c[0]*q + c[1])*q + c[2])*q + c[3])*q + c[4])*q + c[5]) / \
                    ((((d[0]*q + d[1])*q + d[2])*q + d[3])*q + 1)

    def get_parametric_var(self, confidence: float) -> float:
        """
        Calculate parametric VaR using Student-t distribution.

        VaR_α = -μ - σ × √((ν-2)/ν) × t_ν^{-1}(1-α)

        Args:
            confidence: Confidence level (e.g., 0.95, 0.99)

        Returns:
            VaR value (negative number representing loss)
        """
        if not 0 < confidence < 1:
            raise ValueError("Confidence must be in (0, 1)")

        sigma = self._get_sigma()
        mu = self._mean_ema

        # Student-t scale adjustment: √((ν-2)/ν)
        scale = math.sqrt((self._nu - 2) / self._nu) if self._nu > 2 else 1.0

        # Quantile at (1 - confidence) for left tail
        t_q = self._t_quantile(confidence)

        # VaR = -μ - σ × scale × t_quantile
        # Note: t_quantile is positive, VaR should be negative for losses
        return -mu - sigma * scale * t_q

    def get_historical_var(self, confidence: float) -> float:
        """
        Calculate historical VaR from empirical distribution.

        VaR_α = -F^{-1}(1-α) = (1-α) quantile of returns

        Args:
            confidence: Confidence level

        Returns:
            VaR value (negative for loss)
        """
        n = len(self._returns)
        if n < 50:
            return self.get_parametric_var(confidence)

        returns_sorted = sorted(self._returns)
        idx = int((1 - confidence) * n)
        idx = max(0, min(idx, n - 1))

        return returns_sorted[idx]

    def get_cornish_fisher_var(self, confidence: float) -> float:
        """
        Calculate Cornish-Fisher adjusted VaR.

        Adjusts normal quantile for skewness and kurtosis:
        z_cf = z + (z² - 1)S/6 + (z³ - 3z)(K-3)/24 - (2z³ - 5z)S²/36

        where S = skewness, K = kurtosis
        """
        sigma = self._get_sigma()
        mu = self._mean_ema
        z = self._normal_quantile(confidence)

        S = self._skew_ema
        K = self._kurt_ema

        # Cornish-Fisher expansion
        z_cf = (
            z +
            (z * z - 1) * S / 6 +
            (z * z * z - 3 * z) * (K - 3) / 24 -
            (2 * z * z * z - 5 * z) * S * S / 36
        )

        return -mu - sigma * z_cf

    def get_expected_shortfall(self, confidence: float) -> float:
        """
        Calculate Expected Shortfall (CVaR) from historical returns.

        ES_α = E[X | X ≤ -VaR_α] = average of returns worse than VaR

        This is the coherent alternative to VaR, satisfying subadditivity.

        Args:
            confidence: Confidence level

        Returns:
            ES value (negative for expected loss)
        """
        n = len(self._returns)
        if n < 50:
            # Parametric ES for Student-t
            return self._get_parametric_es(confidence)

        var = self.get_historical_var(confidence)
        tail_returns = [r for r in self._returns if r <= var]

        if not tail_returns:
            return var

        return sum(tail_returns) / len(tail_returns)

    def _get_parametric_es(self, confidence: float) -> float:
        """
        Parametric ES for Student-t distribution.

        ES_α = -μ + σ × f(t_α) / (1-α) × (ν + t_α²) / (ν - 1)

        where f is the Student-t pdf, t_α is the quantile.
        """
        sigma = self._get_sigma()
        mu = self._mean_ema

        t_q = self._t_quantile(confidence)
        scale = math.sqrt((self._nu - 2) / self._nu) if self._nu > 2 else 1.0

        # Student-t pdf at quantile (simplified)
        # f(t) ∝ (1 + t²/ν)^{-(ν+1)/2}
        f_t = math.pow(1 + t_q * t_q / self._nu, -(self._nu + 1) / 2)

        # Normalization constant approximation
        norm = 1 / math.sqrt(self._nu * math.pi)

        # ES formula
        es_factor = norm * f_t / (1 - confidence) * (self._nu + t_q * t_q) / (self._nu - 1)

        return -mu - sigma * scale * es_factor

    def get_risk_metrics(self) -> RiskMetrics:
        """Get all risk metrics."""
        return RiskMetrics(
            var_95=self.get_parametric_var(0.95),
            var_99=self.get_parametric_var(0.99),
            es_95=self.get_expected_shortfall(0.95),
            es_99=self.get_expected_shortfall(0.99),
            current_return=self._returns[-1] if self._returns else 0.0,
            mean_return=self._mean_ema,
            std_return=self._get_sigma(),
        )

    def get_stats(self) -> dict:
        """Return model statistics for dashboard."""
        metrics = self.get_risk_metrics()
        stats = self._base_stats()
        stats.update({
            "symbol": self.symbol,
            "history_size": len(self._returns),
            "nu": self._nu,
            "garch_linked": self.garch is not None,
            # VaR (as percentages)
            "var_95_pct": round(metrics.var_95 * 100, 3),
            "var_99_pct": round(metrics.var_99 * 100, 3),
            # ES (as percentages)
            "es_95_pct": round(metrics.es_95 * 100, 3),
            "es_99_pct": round(metrics.es_99 * 100, 3),
            # Historical for comparison
            "var_99_hist_pct": round(self.get_historical_var(0.99) * 100, 3),
            # Statistics
            "mean_return_pct": round(metrics.mean_return * 100, 4),
            "std_return_pct": round(metrics.std_return * 100, 4),
            "skewness": round(self._skew_ema, 3),
            "kurtosis": round(self._kurt_ema, 3),
            # Breaches
            "var_breaches": self._var_breach_count,
            "breach_rate": round(self._var_breach_count / max(self._tick_count, 1) * 100, 2),
        })
        return stats


class EVTVaR(BaseModel):
    """
    Extreme Value Theory VaR using Generalized Pareto Distribution.

    For exceedances above threshold u, GPD models the tail:
    G_{ξ,β}(y) = 1 - (1 + ξy/β)^{-1/ξ}

    EVT-based VaR:
    VaR_α = u + (β/ξ) × [(n/N_u × (1-α))^{-ξ} - 1]

    EVT-based ES:
    ES_α = (VaR_α + β - ξu) / (1 - ξ)

    Crypto tail estimates: ξ ≈ 0.25-0.35 (heavy but finite variance)
    """

    __slots__ = (
        'symbol', '_returns', '_threshold_quantile',
        '_xi', '_beta', '_u', '_n_exceedances',
        '_prev_price', '_ema_decay', '_var_ema',
    )

    def __init__(
        self,
        symbol: str,
        threshold_quantile: float = 0.95,
        history_size: int = 2000,
        warmup_ticks: int = 1000,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize EVT-based VaR.

        Args:
            symbol: Asset symbol
            threshold_quantile: Quantile for POT threshold (0.90-0.95)
            history_size: Return history size
            warmup_ticks: Ticks before stable estimates
        """
        super().__init__(f"evt_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol
        self._threshold_quantile = threshold_quantile
        self._returns: deque = deque(maxlen=history_size)
        self._prev_price = 0.0

        # GPD parameters (estimated online)
        self._xi = 0.3   # Shape (crypto typical: 0.25-0.35)
        self._beta = 0.01  # Scale
        self._u = 0.0    # Threshold
        self._n_exceedances = 0

        self._ema_decay = 0.01
        self._var_ema = 0.0

    async def on_tick(self, tick: Tick) -> None:
        """Update EVT model with new tick."""
        if tick.symbol != self.symbol:
            return

        price = float(tick.mid_price())
        if self._prev_price <= 0:
            self._prev_price = price
            return

        ret = math.log(price / self._prev_price)
        self._returns.append(ret)
        self._prev_price = price

        # Update variance for threshold estimation
        self._var_ema = (1 - self._ema_decay) * self._var_ema + self._ema_decay * ret * ret

        self._check_warmup()

        # Periodically re-estimate GPD parameters
        if self._tick_count % 100 == 0 and len(self._returns) >= 500:
            self._estimate_gpd()

    def _estimate_gpd(self) -> None:
        """Estimate GPD parameters using method of moments."""
        returns_sorted = sorted(self._returns)
        n = len(returns_sorted)

        # Threshold at quantile (looking at left tail, so use negative returns)
        neg_returns = sorted([-r for r in self._returns if r < 0])
        if len(neg_returns) < 50:
            return

        threshold_idx = int((1 - self._threshold_quantile) * len(neg_returns))
        self._u = neg_returns[threshold_idx] if threshold_idx < len(neg_returns) else 0

        # Exceedances
        exceedances = [x - self._u for x in neg_returns if x > self._u]
        self._n_exceedances = len(exceedances)

        if self._n_exceedances < 20:
            return

        # Method of moments estimation
        mean_exc = sum(exceedances) / len(exceedances)
        var_exc = sum((x - mean_exc) ** 2 for x in exceedances) / len(exceedances)

        if mean_exc <= 0 or var_exc <= 0:
            return

        # MoM estimators for GPD
        # β = mean × (1 + ξ)
        # ξ = (1/2) × (mean²/var - 1)
        self._xi = 0.5 * (mean_exc * mean_exc / var_exc - 1)
        self._xi = max(-0.5, min(self._xi, 0.5))  # Bound for stability

        self._beta = mean_exc * (1 + self._xi)
        self._beta = max(self._beta, 1e-8)

    def get_evt_var(self, confidence: float) -> float:
        """
        Calculate EVT-based VaR.

        VaR_α = u + (β/ξ) × [(n/N_u × (1-α))^{-ξ} - 1]
        """
        if self._n_exceedances < 20 or self._xi == 0:
            # Fallback to simple quantile
            if len(self._returns) < 100:
                return -math.sqrt(self._var_ema) * 2.33
            returns_sorted = sorted(self._returns)
            idx = int((1 - confidence) * len(returns_sorted))
            return returns_sorted[max(0, idx)]

        n = len(self._returns)
        N_u = self._n_exceedances

        # EVT VaR formula
        ratio = (n / N_u) * (1 - confidence)

        if self._xi != 0:
            var_exc = self._u + (self._beta / self._xi) * (math.pow(ratio, -self._xi) - 1)
        else:
            var_exc = self._u - self._beta * math.log(ratio)

        return -var_exc  # Return as negative (loss)

    def get_evt_es(self, confidence: float) -> float:
        """
        Calculate EVT-based Expected Shortfall.

        ES_α = (VaR_α + β - ξu) / (1 - ξ)
        """
        var = abs(self.get_evt_var(confidence))

        if self._xi >= 1:
            return -var * 1.5  # Fallback

        es = (var + self._beta - self._xi * self._u) / (1 - self._xi)
        return -es

    def get_stats(self) -> dict:
        """Return model statistics."""
        stats = self._base_stats()
        stats.update({
            "symbol": self.symbol,
            "xi": round(self._xi, 4),
            "beta": round(self._beta, 6),
            "threshold_u": round(self._u, 6),
            "n_exceedances": self._n_exceedances,
            "var_99_evt_pct": round(self.get_evt_var(0.99) * 100, 3),
            "es_99_evt_pct": round(self.get_evt_es(0.99) * 100, 3),
            "tail_index": round(1 / self._xi if self._xi > 0 else float('inf'), 2),
        })
        return stats
