"""
GARCH Family Volatility Models.

Implements GARCH(1,1) and EGARCH for cryptocurrency volatility forecasting.

Mathematical Foundation:
- GARCH(1,1): σₜ² = ω + α×εₜ₋₁² + β×σₜ₋₁²
- Unconditional variance: σ̄² = ω/(1 - α - β)
- Persistence: α + β (crypto: ~0.99, near-integrated)

EGARCH (Nelson 1991):
- log(σₜ²) = ω + α(|zₜ₋₁| - E|zₜ₋₁|) + γzₜ₋₁ + β×log(σₜ₋₁²)
- γ > 0 for crypto (inverse leverage effect)

Crypto-specific considerations:
- α: 0.09-0.37 (higher shock sensitivity than equities)
- β: 0.7-0.9 (high persistence)
- α + β: 0.85-0.99 (near-integrated)
- No overnight gaps (24/7 trading)
"""
import math
from collections import deque
from dataclasses import dataclass

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


@dataclass(slots=True)
class GARCHState:
    """Current state of GARCH model for dashboard display."""
    sigma_sq: float          # Current conditional variance σₜ²
    sigma: float             # Current conditional volatility σₜ
    sigma_annualized: float  # Annualized volatility
    epsilon_sq: float        # Last squared residual εₜ₋₁²
    long_run_var: float      # ω/(1-α-β)
    persistence: float       # α + β
    half_life_ticks: float   # ln(2)/ln(α+β) - shock decay


class GARCH11(BaseModel):
    """
    GARCH(1,1) model with online variance targeting.

    Model: σₜ² = ω + α×εₜ₋₁² + β×σₜ₋₁²

    Uses variance targeting for stable estimation:
    ω = σ̄²(1 - α - β) where σ̄² is sample variance

    Emits VOL_SPIKE signals when current vol exceeds 3× long-run average.

    Attributes:
        symbol: Asset symbol to track
        _alpha: ARCH coefficient (shock sensitivity)
        _beta: GARCH coefficient (persistence)
        _omega: Intercept (calibrated via variance targeting)
        _sigma_sq: Current conditional variance
        _epsilon_sq: Last squared innovation
    """

    __slots__ = (
        'symbol', '_alpha', '_beta', '_omega',
        '_sigma_sq', '_epsilon_sq', '_prev_price',
        '_return_sq_ema', '_return_mean', '_ema_decay',
        '_long_run_var', '_vol_spike_threshold',
    )

    # Annualization factor: per-tick to annual
    # Assuming ~100ms ticks, 315,360,000 ticks per year
    TICKS_PER_YEAR = 365 * 24 * 60 * 60 * 10  # 10 ticks/sec

    def __init__(
        self,
        symbol: str,
        alpha: float = 0.05,
        beta: float = 0.94,
        warmup_ticks: int = 200,
        vol_spike_threshold: float = 3.0,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize GARCH(1,1) model.

        Args:
            symbol: Asset symbol to track (e.g., "BTC")
            alpha: ARCH coefficient α (default 0.05, crypto typical)
            beta: GARCH coefficient β (default 0.94, high persistence)
            warmup_ticks: Ticks before emitting signals
            vol_spike_threshold: Multiple of long-run vol for spike signal
            on_signal: Async callback for signals
        """
        super().__init__(f"garch_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol

        # GARCH parameters (validated)
        if not 0 < alpha < 1:
            raise ValueError(f"alpha must be in (0, 1), got {alpha}")
        if not 0 < beta < 1:
            raise ValueError(f"beta must be in (0, 1), got {beta}")
        if alpha + beta >= 1:
            raise ValueError(f"alpha + beta must be < 1 for stationarity, got {alpha + beta}")

        self._alpha = alpha
        self._beta = beta
        self._omega = 0.0  # Calibrated via variance targeting

        # State variables
        self._sigma_sq = 0.0
        self._epsilon_sq = 0.0
        self._prev_price = 0.0

        # EMA for variance targeting (online estimation of σ̄²)
        self._ema_decay = 2.0 / (warmup_ticks + 1)
        self._return_sq_ema = 0.0
        self._return_mean = 0.0
        self._long_run_var = 0.0

        self._vol_spike_threshold = vol_spike_threshold

    async def on_tick(self, tick: Tick) -> None:
        """
        Update GARCH state with new tick.

        Implements the recursion:
        σₜ² = ω + α×εₜ₋₁² + β×σₜ₋₁²
        """
        if tick.symbol != self.symbol:
            return

        price = float(tick.mid_price())

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

        # Calculate log return
        ret = math.log(price / self._prev_price)
        self._prev_price = price

        # Update return statistics via EMA
        self._update_return_stats(ret)

        # Update GARCH variance
        self._update_variance(ret)

        # Check warmup
        self._check_warmup()

        # Check for vol spike signal
        if self._is_ready and self._long_run_var > 0:
            ratio = self._sigma_sq / self._long_run_var
            if ratio > self._vol_spike_threshold ** 2:
                await self._emit_vol_spike(ratio)

    def _update_return_stats(self, ret: float) -> None:
        """Update EMA statistics for variance targeting."""
        # EMA for mean
        self._return_mean = (
            self._ema_decay * ret +
            (1 - self._ema_decay) * self._return_mean
        )

        # EMA for squared returns (variance proxy)
        ret_sq = ret * ret
        self._return_sq_ema = (
            self._ema_decay * ret_sq +
            (1 - self._ema_decay) * self._return_sq_ema
        )

        # Long-run variance estimate
        self._long_run_var = self._return_sq_ema - self._return_mean ** 2
        self._long_run_var = max(self._long_run_var, 1e-12)

        # Variance targeting: ω = σ̄²(1 - α - β)
        self._omega = self._long_run_var * (1 - self._alpha - self._beta)

    def _update_variance(self, ret: float) -> None:
        """Update conditional variance using GARCH(1,1) recursion."""
        # Innovation (residual) - using demeaned return
        epsilon = ret - self._return_mean

        if self._tick_count == 0:
            # Initialize with sample variance
            self._sigma_sq = self._long_run_var if self._long_run_var > 0 else ret * ret
            self._epsilon_sq = epsilon * epsilon
        else:
            # GARCH(1,1) recursion: σₜ² = ω + α×εₜ₋₁² + β×σₜ₋₁²
            self._sigma_sq = (
                self._omega +
                self._alpha * self._epsilon_sq +
                self._beta * self._sigma_sq
            )

            # Store current squared innovation for next iteration
            self._epsilon_sq = epsilon * epsilon

        # Ensure non-negative (numerical stability)
        self._sigma_sq = max(self._sigma_sq, 1e-16)

    async def _emit_vol_spike(self, ratio: float) -> None:
        """Emit volatility spike signal."""
        sigma = math.sqrt(self._sigma_sq)
        sigma_ann = sigma * math.sqrt(self.TICKS_PER_YEAR)

        await self._emit_signal(
            signal_type=SignalType.VOL_SPIKE,
            symbol=self.symbol,
            confidence=min((ratio - self._vol_spike_threshold) / self._vol_spike_threshold, 1.0),
            metadata={
                "sigma": sigma,
                "sigma_annualized": sigma_ann,
                "ratio_to_long_run": ratio,
                "long_run_sigma": math.sqrt(self._long_run_var),
            },
        )

    def get_sigma(self) -> float:
        """Get current conditional volatility σₜ."""
        return math.sqrt(self._sigma_sq)

    def get_sigma_annualized(self) -> float:
        """Get annualized conditional volatility."""
        return math.sqrt(self._sigma_sq * self.TICKS_PER_YEAR)

    def get_state(self) -> GARCHState:
        """Get full model state."""
        sigma = math.sqrt(self._sigma_sq)
        persistence = self._alpha + self._beta

        # Half-life of volatility shocks: ln(2) / |ln(α+β)|
        if persistence > 0 and persistence < 1:
            half_life = math.log(2) / abs(math.log(persistence))
        else:
            half_life = float('inf')

        return GARCHState(
            sigma_sq=self._sigma_sq,
            sigma=sigma,
            sigma_annualized=sigma * math.sqrt(self.TICKS_PER_YEAR),
            epsilon_sq=self._epsilon_sq,
            long_run_var=self._long_run_var,
            persistence=persistence,
            half_life_ticks=half_life,
        )

    def get_stats(self) -> dict:
        """Return model statistics for dashboard."""
        state = self.get_state()
        stats = self._base_stats()
        stats.update({
            "symbol": self.symbol,
            "sigma": round(state.sigma, 8),
            "sigma_pct": round(state.sigma * 100, 4),
            "sigma_annualized": round(state.sigma_annualized, 4),
            "sigma_annualized_pct": round(state.sigma_annualized * 100, 2),
            "alpha": self._alpha,
            "beta": self._beta,
            "omega": self._omega,
            "persistence": round(state.persistence, 4),
            "half_life_ticks": round(state.half_life_ticks, 1),
            "long_run_sigma": round(math.sqrt(self._long_run_var), 8),
            "long_run_sigma_pct": round(math.sqrt(self._long_run_var) * 100, 4),
        })
        return stats


class EGARCH(BaseModel):
    """
    EGARCH (Exponential GARCH) model with asymmetric response.

    Model: log(σₜ²) = ω + α(|zₜ₋₁| - E|zₜ₋₁|) + γzₜ₋₁ + β×log(σₜ₋₁²)

    where zₜ = εₜ/σₜ is the standardized residual.

    Key advantages over GARCH:
    - No non-negativity constraints (log domain)
    - Asymmetric response via γ parameter
    - γ > 0 for crypto (inverse leverage effect)
    - γ < 0 for equities (traditional leverage effect)

    For crypto: positive shocks increase volatility due to
    speculative buying, opposite to equity markets.
    """

    __slots__ = (
        'symbol', '_alpha', '_beta', '_gamma', '_omega',
        '_log_sigma_sq', '_prev_price', '_prev_z',
        '_return_sq_ema', '_return_mean', '_ema_decay',
    )

    TICKS_PER_YEAR = 365 * 24 * 60 * 60 * 10

    # E[|Z|] for standard normal = sqrt(2/π)
    E_ABS_Z = math.sqrt(2.0 / math.pi)

    def __init__(
        self,
        symbol: str,
        alpha: float = 0.10,
        beta: float = 0.90,
        gamma: float = 0.05,  # Positive for crypto inverse leverage
        warmup_ticks: int = 200,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize EGARCH model.

        Args:
            symbol: Asset symbol to track
            alpha: Magnitude effect coefficient
            beta: Persistence coefficient
            gamma: Asymmetry coefficient (γ > 0 for crypto)
            warmup_ticks: Ticks before emitting signals
            on_signal: Async callback for signals
        """
        super().__init__(f"egarch_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol

        # Parameters (no positivity constraints in EGARCH)
        if not 0 < abs(beta) < 1:
            raise ValueError(f"|beta| must be < 1 for stationarity, got {beta}")

        self._alpha = alpha
        self._beta = beta
        self._gamma = gamma  # Can be positive or negative
        self._omega = 0.0  # Calibrated from data

        # State (log domain for numerical stability)
        self._log_sigma_sq = 0.0
        self._prev_price = 0.0
        self._prev_z = 0.0  # Previous standardized residual

        # EMA for targeting
        self._ema_decay = 2.0 / (warmup_ticks + 1)
        self._return_sq_ema = 0.0
        self._return_mean = 0.0

    async def on_tick(self, tick: Tick) -> None:
        """Update EGARCH state 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._prev_price = price

        # Update statistics
        self._update_return_stats(ret)

        # Update EGARCH
        self._update_variance(ret)

        self._check_warmup()

    def _update_return_stats(self, ret: float) -> None:
        """Update EMA statistics."""
        self._return_mean = (
            self._ema_decay * ret +
            (1 - self._ema_decay) * self._return_mean
        )

        self._return_sq_ema = (
            self._ema_decay * ret * ret +
            (1 - self._ema_decay) * self._return_sq_ema
        )

        # Estimate unconditional log-variance for targeting
        var = max(self._return_sq_ema - self._return_mean ** 2, 1e-12)
        # ω = (1-β) × E[log(σ²)] approximately
        self._omega = (1 - self._beta) * math.log(var)

    def _update_variance(self, ret: float) -> None:
        """
        Update log-variance using EGARCH recursion.

        log(σₜ²) = ω + α(|zₜ₋₁| - E|zₜ₋₁|) + γzₜ₋₁ + β×log(σₜ₋₁²)
        """
        # Current sigma from previous state
        sigma_sq = math.exp(self._log_sigma_sq) if self._log_sigma_sq > -30 else 1e-12
        sigma = math.sqrt(sigma_sq)

        # Innovation
        epsilon = ret - self._return_mean

        # Standardized residual
        z = epsilon / sigma if sigma > 1e-12 else 0.0

        if self._tick_count == 0:
            # Initialize
            var = max(self._return_sq_ema - self._return_mean ** 2, 1e-12)
            self._log_sigma_sq = math.log(var)
        else:
            # EGARCH recursion
            # log(σₜ²) = ω + α(|zₜ₋₁| - E|zₜ₋₁|) + γzₜ₋₁ + β×log(σₜ₋₁²)
            self._log_sigma_sq = (
                self._omega +
                self._alpha * (abs(self._prev_z) - self.E_ABS_Z) +
                self._gamma * self._prev_z +
                self._beta * self._log_sigma_sq
            )

        # Store for next iteration
        self._prev_z = z

        # Clamp for numerical stability
        self._log_sigma_sq = max(min(self._log_sigma_sq, 10), -30)

    def get_sigma(self) -> float:
        """Get current conditional volatility."""
        sigma_sq = math.exp(self._log_sigma_sq)
        return math.sqrt(sigma_sq)

    def get_sigma_annualized(self) -> float:
        """Get annualized volatility."""
        return self.get_sigma() * math.sqrt(self.TICKS_PER_YEAR)

    def get_stats(self) -> dict:
        """Return model statistics."""
        sigma = self.get_sigma()
        sigma_ann = sigma * math.sqrt(self.TICKS_PER_YEAR)

        stats = self._base_stats()
        stats.update({
            "symbol": self.symbol,
            "sigma": round(sigma, 8),
            "sigma_pct": round(sigma * 100, 4),
            "sigma_annualized": round(sigma_ann, 4),
            "sigma_annualized_pct": round(sigma_ann * 100, 2),
            "log_sigma_sq": round(self._log_sigma_sq, 6),
            "alpha": self._alpha,
            "beta": self._beta,
            "gamma": self._gamma,
            "gamma_interpretation": "inverse_leverage" if self._gamma > 0 else "leverage",
        })
        return stats
