"""
Jump Detection Models for Cryptocurrency Price Dynamics.

Implements jump tests and models:
- Lee-Mykland (2008): Bipower variation based test
- Barndorff-Nielsen-Shephard: Realized bipower variation

Mathematical Foundation:
Test statistic: L(i) = rₜ,ᵢ / σ̂ₜ,ᵢ
where σ̂² = (π/2) × (K-2)⁻¹ × Σⱼ |rⱼ| × |rⱼ₋₁| (bipower variation)

Rejection region: (|L(i)| - Cₙ) / Sₙ > -ln(-ln(1-α))
where:
- Cₙ = √(2 ln n) / 0.7979 - (ln π + ln(ln n)) / (1.5958 √(2 ln n))
- Sₙ = 1 / (1.5958 √(2 ln n))

Crypto jump statistics:
- ~3.5 jumps per day on 1-minute data
- Jump intensity: 9.89%-16.85% quarterly variation
- Mean jump size: μⱼ ≈ 0.124
"""
import math
from collections import deque
from dataclasses import dataclass

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


@dataclass(slots=True)
class JumpEvent:
    """Detected jump event."""
    timestamp_ms: int
    symbol: str
    return_value: float       # Return that triggered the jump
    test_statistic: float     # L(i) value
    threshold: float          # Critical value
    bipower_sigma: float      # σ̂ from bipower variation
    is_positive: bool         # Jump direction
    severity: float           # How far above threshold


class LeeMyklandTest(BaseModel):
    """
    Lee-Mykland (2008) jump detection test.

    Uses bipower variation to estimate continuous volatility,
    then tests if individual returns are too large to be
    explained by diffusion alone.

    Bipower variation is robust to jumps:
    BPV = (π/2) × Σ|rᵢ| × |rᵢ₋₁| → ∫σ²dt (under diffusion)

    Whereas realized variance includes jumps:
    RV = Σrᵢ² → ∫σ²dt + Σ(jump sizes)²

    So RV - BPV estimates jump contribution to variance.
    """

    __slots__ = (
        'symbol', '_window_size', '_significance',
        '_returns', '_abs_returns',
        '_prev_price', '_sum_bpv',
        '_jump_count', '_jumps_today', '_last_jump_time',
    )

    def __init__(
        self,
        symbol: str,
        window_size: int = 100,
        significance: float = 0.001,
        warmup_ticks: int = 200,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize Lee-Mykland jump test.

        Args:
            symbol: Asset symbol to track
            window_size: Window for bipower variation (K)
            significance: Test significance level α (default 0.001)
            warmup_ticks: Ticks before testing
            on_signal: Async callback for jump signals
        """
        super().__init__(f"jump_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol
        self._window_size = window_size
        self._significance = significance

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

        # Running sum for bipower variation
        self._sum_bpv = 0.0

        self._prev_price = 0.0
        self._jump_count = 0
        self._jumps_today = 0
        self._last_jump_time = 0

    async def on_tick(self, tick: Tick) -> None:
        """Update jump test 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)
        abs_ret = abs(ret)
        self._prev_price = price

        # Update bipower variation sum (maintain O(1))
        if len(self._abs_returns) == self._window_size:
            # Remove old contribution
            if len(self._abs_returns) >= 2:
                old_abs = self._abs_returns[0]
                old_abs_prev = self._abs_returns[1]
                self._sum_bpv -= old_abs * old_abs_prev

        # Add new return
        self._returns.append(ret)

        # Update BPV sum
        if len(self._abs_returns) > 0:
            self._sum_bpv += abs_ret * self._abs_returns[-1]

        self._abs_returns.append(abs_ret)

        self._check_warmup()

        # Perform jump test
        if self._is_ready and len(self._abs_returns) >= 30:
            await self._test_for_jump(ret, tick.local_ts)

    async def _test_for_jump(self, ret: float, timestamp: int) -> None:
        """
        Perform Lee-Mykland jump test on current return.

        Test: L(i) = rᵢ / σ̂ᵢ where σ̂ is from bipower variation
        Reject if (|L| - Cₙ) / Sₙ > critical value
        """
        # Estimate σ using bipower variation
        sigma_bpv = self._get_bipower_sigma()
        if sigma_bpv < 1e-10:
            return

        # Test statistic
        L = ret / sigma_bpv

        # Gumbel parameters for extreme value distribution
        n = len(self._abs_returns)
        C_n, S_n = self._gumbel_params(n)

        # Standardized test statistic
        test_stat = (abs(L) - C_n) / S_n

        # Critical value from Gumbel distribution
        # P(max < x) = exp(-exp(-x))
        # For α, critical value = -ln(-ln(1-α))
        critical = -math.log(-math.log(1 - self._significance))

        # Test for jump
        if test_stat > critical:
            self._jump_count += 1
            self._jumps_today += 1
            self._last_jump_time = timestamp

            jump = JumpEvent(
                timestamp_ms=timestamp,
                symbol=self.symbol,
                return_value=ret,
                test_statistic=L,
                threshold=C_n + S_n * critical,
                bipower_sigma=sigma_bpv,
                is_positive=ret > 0,
                severity=(test_stat - critical) / critical,
            )

            await self._emit_jump_signal(jump)

    def _get_bipower_sigma(self) -> float:
        """
        Estimate volatility using bipower variation.

        BPV = (π/2) × (K-2)⁻¹ × Σⱼ |rⱼ| × |rⱼ₋₁|
        σ̂ = √(BPV)

        Bipower variation is robust to jumps.
        """
        n = len(self._abs_returns)
        if n < 3:
            return 0.0

        # BPV = μ₁⁻² × (K-2)⁻¹ × Σ|rⱼ||rⱼ₋₁|
        # where μ₁ = √(2/π) ≈ 0.7979
        # So BPV = (π/2) × sum / (K-2)

        bpv = (math.pi / 2) * self._sum_bpv / (n - 2)
        return math.sqrt(max(bpv, 1e-16))

    def _gumbel_params(self, n: int) -> tuple[float, float]:
        """
        Compute Gumbel distribution parameters for sample size n.

        Cₙ = √(2 ln n) / c - (ln π + ln(ln n)) / (2c √(2 ln n))
        Sₙ = 1 / (c √(2 ln n))

        where c = √(2/π) ≈ 0.7979
        """
        if n < 2:
            return 0.0, 1.0

        c = math.sqrt(2 / math.pi)  # ≈ 0.7979
        sqrt_2ln_n = math.sqrt(2 * math.log(n))

        C_n = sqrt_2ln_n / c - (math.log(math.pi) + math.log(math.log(n))) / (2 * c * sqrt_2ln_n)
        S_n = 1 / (c * sqrt_2ln_n)

        return C_n, S_n

    async def _emit_jump_signal(self, jump: JumpEvent) -> None:
        """Emit jump detection signal."""
        await self._emit_signal(
            signal_type=SignalType.JUMP_DETECTED,
            symbol=self.symbol,
            confidence=min(jump.severity, 1.0),
            metadata={
                "return_pct": round(jump.return_value * 100, 4),
                "test_statistic": round(jump.test_statistic, 3),
                "threshold": round(jump.threshold, 3),
                "bipower_sigma": round(jump.bipower_sigma, 6),
                "direction": "up" if jump.is_positive else "down",
                "severity": round(jump.severity, 3),
                "jump_number": self._jump_count,
            },
            direction="long" if jump.is_positive else "short",
        )

    def get_jump_intensity(self) -> float:
        """
        Estimate jump intensity (jumps per unit time).

        Returns approximate jumps per 1000 ticks.
        """
        if self._tick_count < 100:
            return 0.0
        return self._jump_count / self._tick_count * 1000

    def get_bipower_variation(self) -> float:
        """Get current bipower variation estimate."""
        n = len(self._abs_returns)
        if n < 3:
            return 0.0
        return (math.pi / 2) * self._sum_bpv / (n - 2)

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

        JV / RV = (RV - BPV) / RV
        """
        bpv = self.get_bipower_variation()
        if len(self._returns) < 10:
            return 0.0

        rv = sum(r * r for r in self._returns)
        if rv < 1e-16:
            return 0.0

        return max(0, (rv - bpv) / rv)

    def get_stats(self) -> dict:
        """Return model statistics for dashboard."""
        stats = self._base_stats()

        bpv_sigma = self._get_bipower_sigma()
        C_n, S_n = self._gumbel_params(len(self._abs_returns))

        stats.update({
            "symbol": self.symbol,
            "window_size": self._window_size,
            "significance": self._significance,
            "samples": len(self._returns),
            # Jump counts
            "total_jumps": self._jump_count,
            "jump_intensity_per_1k": round(self.get_jump_intensity(), 2),
            # Volatility estimates
            "bipower_sigma_pct": round(bpv_sigma * 100, 4),
            "bipower_variation": round(self.get_bipower_variation(), 10),
            "jump_variance_ratio": round(self.get_jump_variance_ratio(), 4),
            # Test parameters
            "gumbel_C_n": round(C_n, 3),
            "gumbel_S_n": round(S_n, 4),
            "current_threshold": round(C_n + S_n * (-math.log(-math.log(1 - self._significance))), 3),
        })
        return stats


class BNSJumpTest(BaseModel):
    """
    Barndorff-Nielsen-Shephard (2006) jump test.

    Tests for jumps using the ratio statistic:
    z_BNS = (RV - BPV) / √(Θ × max(TQ - BPV², 0))

    where:
    - RV = realized variance
    - BPV = bipower variation
    - TQ = tripower quarticity (for consistent variance estimation)
    - Θ = (π²/4 + π - 5) ≈ 0.61

    Under no jumps: z_BNS → N(0,1)
    """

    __slots__ = (
        'symbol', '_window_size',
        '_returns', '_abs_returns',
        '_sum_rv', '_sum_bpv', '_sum_tq',
        '_prev_price',
    )

    # Constant for variance estimation
    THETA = math.pi ** 2 / 4 + math.pi - 5  # ≈ 0.61

    def __init__(
        self,
        symbol: str,
        window_size: int = 100,
        warmup_ticks: int = 200,
        on_signal: SignalCallback | None = None,
    ):
        """Initialize BNS jump test."""
        super().__init__(f"bns_{symbol}", warmup_ticks, on_signal)
        self.symbol = symbol
        self._window_size = window_size

        self._returns: deque = deque(maxlen=window_size)
        self._abs_returns: deque = deque(maxlen=window_size)

        self._sum_rv = 0.0
        self._sum_bpv = 0.0
        self._sum_tq = 0.0  # Tripower quarticity

        self._prev_price = 0.0

    async def on_tick(self, tick: Tick) -> None:
        """Update BNS test 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)
        abs_ret = abs(ret)
        self._prev_price = price

        # Update running sums
        self._update_sums(ret, abs_ret)
        self._check_warmup()

    def _update_sums(self, ret: float, abs_ret: float) -> None:
        """Update running sums for RV, BPV, TQ."""
        # Remove old values if window full
        if len(self._returns) == self._window_size:
            old_ret = self._returns[0]
            self._sum_rv -= old_ret * old_ret

            if len(self._abs_returns) >= 2:
                old_abs = self._abs_returns[0]
                old_abs_prev = self._abs_returns[1]
                self._sum_bpv -= old_abs * old_abs_prev

            if len(self._abs_returns) >= 3:
                # Tripower: |r_i|^{4/3} × |r_{i-1}|^{4/3} × |r_{i-2}|^{4/3}
                old_abs = self._abs_returns[0]
                old_abs_1 = self._abs_returns[1]
                old_abs_2 = self._abs_returns[2]
                self._sum_tq -= (
                    math.pow(old_abs, 4/3) *
                    math.pow(old_abs_1, 4/3) *
                    math.pow(old_abs_2, 4/3)
                )

        # Add new return
        self._returns.append(ret)
        self._sum_rv += ret * ret

        # Update BPV
        if len(self._abs_returns) > 0:
            self._sum_bpv += abs_ret * self._abs_returns[-1]

        # Update TQ
        if len(self._abs_returns) >= 2:
            self._sum_tq += (
                math.pow(abs_ret, 4/3) *
                math.pow(self._abs_returns[-1], 4/3) *
                math.pow(self._abs_returns[-2], 4/3)
            )

        self._abs_returns.append(abs_ret)

    def get_z_statistic(self) -> float:
        """
        Compute BNS test statistic.

        z = (RV - BPV) / √(Θ × max(TQ - BPV², 0))
        """
        n = len(self._returns)
        if n < 10:
            return 0.0

        rv = self._sum_rv
        bpv = (math.pi / 2) * self._sum_bpv / max(n - 1, 1)

        # Tripower quarticity scaling
        mu_43 = math.pow(2, 2/3) * math.gamma(7/6) / math.gamma(0.5)
        tq = self._sum_tq * n / (n - 2) / math.pow(mu_43, 3)

        # Variance of RV - BPV
        var_term = max(tq - bpv * bpv, 0)
        variance = self.THETA * var_term

        if variance < 1e-16:
            return 0.0

        z = (rv - bpv) / math.sqrt(variance)
        return z

    def has_jumps(self, significance: float = 0.05) -> bool:
        """Test if there are significant jumps at given level."""
        z = self.get_z_statistic()
        # One-sided test: jumps increase RV relative to BPV
        # Critical value for 5%: 1.645
        critical = 1.645 if significance == 0.05 else 2.326  # 1%
        return z > critical

    def get_stats(self) -> dict:
        """Return model statistics."""
        stats = self._base_stats()
        z = self.get_z_statistic()

        stats.update({
            "symbol": self.symbol,
            "z_statistic": round(z, 3),
            "has_jumps_5pct": self.has_jumps(0.05),
            "has_jumps_1pct": self.has_jumps(0.01),
            "rv": round(self._sum_rv, 10),
            "bpv": round((math.pi / 2) * self._sum_bpv / max(len(self._returns) - 1, 1), 10),
        })
        return stats
