"""
Realized Volatility Estimators.

Implements high-frequency volatility estimators:
- Realized Variance (RV): Sum of squared returns
- Bipower Variation (BPV): Jump-robust estimator
- Parkinson: High-low range estimator
- Garman-Klass: OHLC-based estimator

Mathematical foundations:
- RV_t = Σᵢ rₜ,ᵢ² converges to integrated variance
- BPV = (π/2) × Σᵢ |rᵢ| × |rᵢ₋₁| is robust to jumps
- Jump variance: JV = max(RV - BPV, 0)

Used for:
- GARCH model validation
- Jump detection (comparing RV vs BPV)
- Volatility forecasting evaluation
"""
import math
from collections import deque
from dataclasses import dataclass

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


@dataclass(slots=True)
class VolatilityEstimates:
    """Container for multiple volatility estimates."""
    realized_variance: float      # RV = Σr²
    bipower_variation: float      # BPV (jump-robust)
    jump_variance: float          # max(RV - BPV, 0)
    parkinson_vol: float          # High-low based
    realized_vol: float           # √RV
    realized_vol_annualized: float


class RealizedVolatility(BaseModel):
    """
    Real-time realized volatility estimator.

    Computes multiple volatility measures in O(1) per tick:
    - Realized Variance: RV = Σrₜ² (simple sum of squared returns)
    - Bipower Variation: BPV = (π/2) × Σ|rᵢ||rᵢ₋₁| (jump-robust)
    - Jump Variance: JV = max(RV - BPV, 0)

    Uses rolling windows for continuous estimation.

    The bipower variation is crucial for jump detection:
    - Under no jumps: BPV ≈ RV
    - With jumps: BPV < RV, difference indicates jump contribution
    """

    __slots__ = (
        'symbol', '_window_size',
        '_returns', '_abs_returns', '_prices',
        '_sum_r_sq', '_sum_bpv',
        '_prev_price', '_high', '_low',
    )

    TICKS_PER_YEAR = 365 * 24 * 60 * 60 * 10

    def __init__(
        self,
        symbol: str,
        window_size: int = 100,
        warmup_ticks: int = 100,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize realized volatility estimator.

        Args:
            symbol: Asset symbol to track
            window_size: Rolling window size for RV calculation
            warmup_ticks: Ticks before stable estimates
            on_signal: Async callback
        """
        super().__init__(f"rv_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol
        self._window_size = window_size

        # Rolling buffers
        self._returns: deque = deque(maxlen=window_size)
        self._abs_returns: deque = deque(maxlen=window_size)
        self._prices: deque = deque(maxlen=window_size)

        # Running sums for O(1) updates
        self._sum_r_sq = 0.0  # Σr²
        self._sum_bpv = 0.0   # Σ|rᵢ||rᵢ₋₁|

        self._prev_price = 0.0
        self._high = 0.0
        self._low = float('inf')

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

        price = float(tick.mid_price())
        self._prices.append(price)

        # Track high/low for Parkinson
        self._high = max(self._high, price)
        self._low = min(self._low, price)

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

        # Log return
        ret = math.log(price / self._prev_price)
        abs_ret = abs(ret)

        # Update running sums (subtract old value if window full)
        if len(self._returns) == self._window_size:
            old_ret = self._returns[0]
            self._sum_r_sq -= old_ret * old_ret

            # Update BPV sum (remove old contribution)
            if len(self._abs_returns) >= 2:
                old_abs = self._abs_returns[0]
                old_abs_prev = self._abs_returns[1] if len(self._abs_returns) > 1 else 0
                self._sum_bpv -= old_abs * old_abs_prev

        # Add new values
        self._returns.append(ret)
        self._sum_r_sq += ret * ret

        # Update BPV: |rᵢ| × |rᵢ₋₁|
        if len(self._abs_returns) > 0:
            prev_abs = self._abs_returns[-1]
            self._sum_bpv += abs_ret * prev_abs

        self._abs_returns.append(abs_ret)
        self._prev_price = price

        self._check_warmup()

    def get_realized_variance(self) -> float:
        """
        Get realized variance RV = Σr².

        This is the sum of squared returns, converging to
        integrated variance under diffusion.
        """
        n = len(self._returns)
        if n == 0:
            return 0.0
        return self._sum_r_sq

    def get_bipower_variation(self) -> float:
        """
        Get bipower variation BPV = (π/2) × Σ|rᵢ||rᵢ₋₁|.

        Jump-robust estimator: converges to integrated variance
        even in presence of jumps (Barndorff-Nielsen & Shephard).
        """
        n = len(self._abs_returns)
        if n < 2:
            return 0.0

        # BPV = μ₁⁻² × Σ|rᵢ||rᵢ₋₁| where μ₁ = √(2/π)
        # Simplifies to (π/2) × Σ|rᵢ||rᵢ₋₁|
        return (math.pi / 2) * self._sum_bpv

    def get_jump_variance(self) -> float:
        """
        Get jump variance JV = max(RV - BPV, 0).

        Estimates the contribution of jumps to total variance.
        """
        rv = self.get_realized_variance()
        bpv = self.get_bipower_variation()
        return max(rv - bpv, 0.0)

    def get_parkinson_vol(self) -> float:
        """
        Get Parkinson volatility estimator.

        σ_P = √(ln(H/L)² / (4×ln(2)))

        More efficient than close-to-close for continuous data.
        """
        if self._high <= self._low or self._low <= 0:
            return 0.0

        log_range = math.log(self._high / self._low)
        return log_range / (2 * math.sqrt(math.log(2)))

    def get_estimates(self) -> VolatilityEstimates:
        """Get all volatility estimates."""
        rv = self.get_realized_variance()
        bpv = self.get_bipower_variation()
        jv = max(rv - bpv, 0.0)
        parkinson = self.get_parkinson_vol()
        rv_vol = math.sqrt(rv) if rv > 0 else 0.0

        return VolatilityEstimates(
            realized_variance=rv,
            bipower_variation=bpv,
            jump_variance=jv,
            parkinson_vol=parkinson,
            realized_vol=rv_vol,
            realized_vol_annualized=rv_vol * math.sqrt(self.TICKS_PER_YEAR),
        )

    def get_jump_ratio(self) -> float:
        """
        Get ratio of jump variance to total variance.

        JR = JV / RV = (RV - BPV) / RV

        Higher values indicate more jump activity.
        """
        rv = self.get_realized_variance()
        if rv < 1e-16:
            return 0.0
        jv = self.get_jump_variance()
        return jv / rv

    def reset_high_low(self) -> None:
        """Reset high/low tracking (e.g., at period boundary)."""
        if len(self._prices) > 0:
            current = self._prices[-1]
            self._high = current
            self._low = current
        else:
            self._high = 0.0
            self._low = float('inf')

    def get_stats(self) -> dict:
        """Return model statistics for dashboard."""
        estimates = self.get_estimates()
        stats = self._base_stats()
        stats.update({
            "symbol": self.symbol,
            "window_size": self._window_size,
            "samples": len(self._returns),
            "realized_variance": round(estimates.realized_variance, 10),
            "realized_vol": round(estimates.realized_vol, 8),
            "realized_vol_pct": round(estimates.realized_vol * 100, 4),
            "realized_vol_annualized": round(estimates.realized_vol_annualized, 4),
            "realized_vol_annualized_pct": round(estimates.realized_vol_annualized * 100, 2),
            "bipower_variation": round(estimates.bipower_variation, 10),
            "jump_variance": round(estimates.jump_variance, 10),
            "jump_ratio": round(self.get_jump_ratio(), 4),
            "parkinson_vol": round(estimates.parkinson_vol, 8),
        })
        return stats
