"""
Lead-Lag Arbitrage Brain.

Implements statistical arbitrage based on BTC leading altcoin price movements.
Uses EMA-based statistics, beta coefficients, and Z-scores to detect predictive
signals when BTC moves but altcoins haven't yet responded.

Mathematical Foundation:
- Ornstein-Uhlenbeck process for mean reversion: dXₜ = μ(θ − Xₜ)dt + σdWₜ
- Rolling Beta: β = Cov(R_alt, R_btc) / Var(R_btc)
- Idiosyncratic volatility: σ_ε = σ_total × √(1 - ρ²)
- Residual Z-score: Z = (R_gap) / σ_ε (NOT σ_total!)
- Half-life: t_0.5 = ln(2) / |λ| where λ is mean-reversion speed

Key Implementation Details:
1. EMA-based statistics (not sliding window) for numerical stability
   - Cov_t = α × dx × dy + (1-α) × Cov_{t-1}
   - No catastrophic cancellation over millions of ticks
2. Idiosyncratic vol for residual normalization
   - Total vol σ_total includes systematic (BTC) risk
   - Only σ_ε (residual vol) is appropriate for gap Z-score
   - With ρ=0.9, σ_ε ≈ 0.43×σ_total; using σ_total loses 60%+ of trades

References:
- arXiv:2501.03171 - High-frequency lead-lag relationships
- arXiv:2403.12180 - Statistical arbitrage with RL
- arXiv:2405.15461 - Market-neutral crypto pair trading
"""
from collections import deque
from dataclasses import dataclass, field
from decimal import Decimal
from enum import Enum
from typing import Callable, Awaitable
import math
import time

from models import Tick
from utils import get_logger


class SignalDirection(Enum):
    """Direction of the predicted move."""
    LONG = "long"    # Altcoin expected to rise
    SHORT = "short"  # Altcoin expected to fall
    NEUTRAL = "neutral"


@dataclass(slots=True)
class LeadLagSignal:
    """Signal generated when lead-lag opportunity detected."""
    timestamp_ms: int
    leader_symbol: str          # "BTC"
    lagger_symbol: str          # e.g., "SOL", "ETH"
    direction: SignalDirection

    # Leader statistics
    leader_return_pct: float    # BTC return that triggered signal
    leader_z_score: float       # Z-score of BTC return

    # Lagger statistics
    lagger_return_pct: float    # Current altcoin return
    expected_return_pct: float  # β × leader_return
    return_gap_pct: float       # expected - actual (opportunity size)

    # Model parameters
    rolling_beta: float         # Current β estimate
    rolling_correlation: float  # Current correlation
    confidence: float           # Signal confidence (0-1)

    # Timing
    max_hold_ms: int            # Suggested max hold time based on half-life

    # Exchange-specific: where to execute the trade
    exchange: str = ""          # Exchange where the lagging altcoin was detected
    entry_price: float = 0.0    # Current price for position entry

    # Volatility stats for logging (Bug Fix #1: was missing, causing zero in logs)
    total_std: float = 0.0           # σ_total - total volatility
    idiosyncratic_std: float = 0.0   # σ_ε - residual volatility after removing BTC risk

    def __repr__(self) -> str:
        exch_str = f"@{self.exchange}" if self.exchange else ""
        return (
            f"LeadLagSignal({self.lagger_symbol}{exch_str} {self.direction.value}: "
            f"BTC Z={self.leader_z_score:.2f}, β={self.rolling_beta:.2f}, "
            f"gap={self.return_gap_pct:.2f}%, conf={self.confidence:.2f})"
        )


class EMAStats:
    """
    Exponential Moving Average (EMA) based statistics.

    Numerically stable O(1) updates without sliding window drift.
    Uses EMA for mean, variance, and covariance calculations.

    Key advantages over Welford sliding window:
    1. No catastrophic cancellation from subtracting large floats
    2. O(1) memory (no deques needed for covariance)
    3. Naturally adapts to regime changes
    4. No floating-point drift over millions of ticks

    Mathematical basis:
    - EMA mean: μ_t = α * x_t + (1-α) * μ_{t-1}
    - EMA variance: σ²_t = α * (x_t - μ_t)² + (1-α) * σ²_{t-1}
    - EMA covariance: Cov_t = α * dx * dy + (1-α) * Cov_{t-1}
      where dx = x_t - μ_x, dy = y_t - μ_y
    """

    __slots__ = (
        'alpha', 'count', 'min_samples',
        '_prev_price', '_current_price', '_latest_return',
        '_mean', '_var',
        '_btc_mean', '_btc_var', '_covar',
        '_btc_log_price',  # For spread calculation
        'returns', 'prices', 'spreads',  # spreads for OU half-life estimation
    )

    def __init__(self, window_size: int = 100):
        """
        Initialize EMA statistics.

        Args:
            window_size: Effective window size. Alpha = 2/(window+1).
                        For window=100, alpha≈0.02 (2% weight on new data).
        """
        # EMA decay factor: α = 2/(N+1) gives similar effective window to SMA
        self.alpha = 2.0 / (window_size + 1)
        self.count = 0
        self.min_samples = window_size // 2

        # Price tracking
        self._prev_price: float = 0.0
        self._current_price: float = 0.0
        self._latest_return: float = 0.0

        # EMA statistics for this asset
        self._mean: float = 0.0
        self._var: float = 0.0

        # EMA statistics for BTC (reference)
        self._btc_mean: float = 0.0
        self._btc_var: float = 0.0

        # EMA covariance with BTC
        self._covar: float = 0.0

        # Track BTC log price for spread calculation
        self._btc_log_price: float = 0.0

        # Keep small deques for half-life estimation and dashboard display
        self.returns: deque = deque(maxlen=50)
        self.prices: deque = deque(maxlen=10)  # For dashboard only

        # CRITICAL: Track the SPREAD (cointegration residual) for OU half-life
        # Spread = ln(P_alt) - β × ln(P_btc)
        # This is what mean-reverts in stat arb, NOT returns
        self.spreads: deque = deque(maxlen=100)

    def update(
        self,
        price: float,
        btc_return: float | None = None,
        btc_log_price: float | None = None,
    ) -> float | None:
        """
        Update statistics with new price using EMA.

        Args:
            price: New price observation
            btc_return: Corresponding BTC return (for covariance)
            btc_log_price: Current ln(P_btc) for spread calculation

        Returns:
            The calculated return, or None if first observation.
        """
        # Track current price for dashboard
        self._current_price = price
        self.prices.append(price)

        # Track BTC log price for spread calculation
        if btc_log_price is not None:
            self._btc_log_price = btc_log_price

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

        # Calculate log return
        ret = math.log(price / self._prev_price)
        self._prev_price = price
        self._latest_return = ret
        self.count += 1

        # Store for half-life estimation (legacy, still used for fallback)
        self.returns.append(ret)

        # Calculate and store the SPREAD (cointegration residual)
        # S_t = ln(P_alt) - β × ln(P_btc)
        # This is the OU process that actually mean-reverts
        if self._btc_log_price > 0 and self.count > self.min_samples:
            log_price = math.log(price)
            spread = log_price - self.beta * self._btc_log_price
            self.spreads.append(spread)

        # Update EMA statistics
        if self.count == 1:
            # Initialize with first observation
            self._mean = ret
            self._var = 0.0
        else:
            # EMA update for mean
            delta = ret - self._mean
            self._mean += self.alpha * delta

            # EMA update for variance: Var_t = (1-α) * (Var_{t-1} + α * δ²)
            # This is the "exponentially weighted" version
            self._var = (1 - self.alpha) * (self._var + self.alpha * delta * delta)

        # Update covariance with BTC if provided
        if btc_return is not None:
            self._update_covariance(ret, btc_return)

        return ret

    def _update_covariance(self, ret: float, btc_ret: float) -> None:
        """
        Update EMA-based covariance with BTC.

        Uses the formula: Cov_t = α * (x - μ_x)(y - μ_y) + (1-α) * Cov_{t-1}
        This is numerically stable and O(1).
        """
        if self.count == 1:
            self._btc_mean = btc_ret
            self._btc_var = 0.0
            self._covar = 0.0
            return

        # EMA update for BTC mean
        btc_delta = btc_ret - self._btc_mean
        self._btc_mean += self.alpha * btc_delta

        # EMA update for BTC variance
        self._btc_var = (1 - self.alpha) * (self._btc_var + self.alpha * btc_delta * btc_delta)

        # EMA update for covariance
        # Use deviations from updated means
        dx = ret - self._mean
        dy = btc_ret - self._btc_mean
        self._covar = self.alpha * dx * dy + (1 - self.alpha) * self._covar

    @property
    def mean(self) -> float:
        """EMA mean of returns."""
        return self._mean

    @property
    def variance(self) -> float:
        """EMA variance of returns."""
        return max(0.0, self._var)

    @property
    def std(self) -> float:
        """EMA standard deviation of returns (total volatility σ_total)."""
        return math.sqrt(self.variance)

    @property
    def btc_variance(self) -> float:
        """EMA variance of BTC returns."""
        return max(0.0, self._btc_var)

    @property
    def covariance(self) -> float:
        """EMA covariance with BTC returns."""
        return self._covar

    @property
    def beta(self) -> float:
        """
        Rolling beta: β = Cov(R_alt, R_btc) / Var(R_btc)

        Measures systematic exposure to BTC.
        """
        if self._btc_var < 1e-12:
            return 1.0
        return self._covar / self._btc_var

    @property
    def correlation(self) -> float:
        """
        Rolling Pearson correlation with BTC: ρ = Cov / (σ_x × σ_y)
        """
        if self._var < 1e-12 or self._btc_var < 1e-12:
            return 0.0
        return self._covar / (math.sqrt(self._var) * math.sqrt(self._btc_var))

    @property
    def idiosyncratic_std(self) -> float:
        """
        Idiosyncratic volatility: σ_ε = σ_total × √(1 - ρ²)

        This is the residual volatility after removing systematic (BTC) risk.
        Critical for proper Z-score calculation in lead-lag signals.

        From: R_alt = β × R_btc + ε
        We have: Var(R_alt) = β² × Var(R_btc) + Var(ε)
        Therefore: σ_ε = σ_total × √(1 - ρ²)
        """
        rho = self.correlation
        rho_squared = min(rho * rho, 0.9999)  # Cap to avoid sqrt of negative
        return self.std * math.sqrt(1 - rho_squared)

    def z_score(self, value: float) -> float:
        """
        Calculate Z-score using total volatility.

        Z = (x - μ) / σ_total
        """
        if self.std < 1e-12:
            return 0.0
        return (value - self._mean) / self.std

    def residual_z_score(self, residual: float) -> float:
        """
        Calculate Z-score for a residual using IDIOSYNCRATIC volatility.

        This is the correct normalization for lead-lag signals.
        Z = residual / σ_ε

        When correlation is high (e.g., 0.9), idiosyncratic vol is ~43% of total.
        Using total vol would understate Z-scores by 2.3x, missing trades.
        """
        sigma_eps = self.idiosyncratic_std
        if sigma_eps < 1e-12:
            return 0.0
        return residual / sigma_eps

    def is_ready(self) -> bool:
        """Check if enough samples for reliable statistics."""
        return self.count >= self.min_samples


# Backwards compatibility alias
RollingStats = EMAStats


@dataclass
class LeadLagConfig:
    """Configuration for the lead-lag model."""

    # Rolling window sizes
    stats_window: int = 100          # Window for rolling mean/std
    beta_window: int = 200           # Window for beta calculation

    # Signal thresholds
    leader_z_threshold: float = 2.0  # Min |Z| for BTC to trigger signal
    lag_z_threshold: float = 1.0     # Max |Z| for altcoin (hasn't moved yet)
    min_correlation: float = 0.5     # Minimum |correlation| to trust beta (HARD STOP)
    min_confidence: float = 0.6      # Minimum confidence to emit signal

    # Return gap threshold (in standard deviations)
    # Signal when: |expected_return - actual_return| > gap_threshold * std
    gap_threshold: float = 1.5

    # Timing
    max_lag_ms: int = 60_000         # Maximum lag window (60 seconds)
    half_life_multiplier: float = 2.0  # Hold for 2x half-life
    default_half_life_ms: int = 30_000  # Default if can't calculate

    # Decay factor for signal confidence over time
    confidence_decay_per_sec: float = 0.05

    # Fee filter (Bug Fix #4): Minimum expected return to cover round-trip fees
    # Typical taker fee: 0.05%-0.10% per side = 0.10%-0.20% round-trip
    # Set to 2x round-trip fees to ensure profitability
    min_expected_return_pct: float = 0.25  # 25 bps minimum (covers ~12.5bps per side)


SignalCallback = Callable[[LeadLagSignal], Awaitable[None]]


class LeadLagBrain:
    """
    Statistical arbitrage engine for lead-lag relationships.

    Monitors BTC (leader) and altcoins (laggers), detecting opportunities
    when BTC makes a significant move that altcoins haven't yet followed.

    The model uses:
    1. Rolling statistics for return normalization
    2. Rolling beta for expected response magnitude
    3. Z-score threshold for signal generation
    4. Half-life estimation for holding period
    """

    def __init__(
        self,
        config: LeadLagConfig | None = None,
        on_signal: SignalCallback | None = None,
        lead_exchange: str = "binance",
    ):
        """
        Initialize the lead-lag brain.

        Args:
            config: Model configuration parameters.
            on_signal: Callback invoked when opportunity detected.
            lead_exchange: Name of the lead exchange (for BTC tracking).
        """
        self.config = config or LeadLagConfig()
        self.on_signal = on_signal
        self.lead_exchange = lead_exchange
        self.logger = get_logger("lead_lag_brain")

        # Rolling statistics for BTC (leader) - single exchange
        self._btc_stats: EMAStats | None = None

        # Rolling statistics per (exchange, symbol) for altcoins (laggers)
        # Key: (exchange, symbol) e.g., ("coinbase", "ETH")
        # This allows us to track how far behind each exchange's altcoin is
        self._lagger_stats: dict[tuple[str, str], EMAStats] = {}

        # Latest BTC return and log price for covariance/spread updates
        self._latest_btc_return: float | None = None
        self._latest_btc_log_price: float = 0.0  # For spread calculation
        self._latest_btc_z: float = 0.0
        self._btc_signal_time_ms: int = 0
        self._btc_return_count: int = 0  # Increment on each new BTC return

        # Track last BTC return count used per (exchange, symbol) to avoid
        # passing the same BTC return multiple times (which kills BTC variance)
        self._last_btc_count_used: dict[tuple[str, str], int] = {}

        # Half-life estimates per symbol (for holding time)
        self._half_lives: dict[str, float] = {}

        # Signal tracking
        self._signals_emitted: int = 0
        # Key: (exchange, symbol) to prevent rapid-fire signals per exchange
        self._last_signal_time: dict[tuple[str, str], int] = {}

        # Static beta overrides from research (fallback values)
        self._research_betas = {
            "ETH": 0.85,
            "SOL": 1.98,
            "XRP": 1.0,
            "DOGE": 1.5,
            "BNB": 1.34,
            "MATIC": 2.07,
            "AVAX": 2.38,
        }

    def _get_btc_stats(self) -> RollingStats:
        """Get or create rolling statistics for BTC (leader)."""
        if self._btc_stats is None:
            self._btc_stats = RollingStats(window_size=self.config.stats_window)
        return self._btc_stats

    def _get_lagger_stats(self, exchange: str, symbol: str) -> RollingStats:
        """Get or create rolling statistics for an altcoin on a specific exchange."""
        key = (exchange, symbol)
        if key not in self._lagger_stats:
            self._lagger_stats[key] = RollingStats(window_size=self.config.stats_window)
        return self._lagger_stats[key]

    async def on_tick(self, tick: Tick) -> None:
        """
        Process an incoming tick.

        - BTC ticks update the leader statistics (from lead exchange only)
        - Altcoin ticks update per-exchange lagger statistics

        This allows detecting when a specific exchange's altcoin hasn't
        moved yet even though BTC (and possibly other exchanges) have.

        Args:
            tick: Normalized tick from any exchange.
        """
        symbol = tick.symbol
        exchange = tick.exchange
        price = float(tick.mid_price())
        now_ms = int(time.time() * 1000)

        if symbol == "BTC":
            # Only track BTC from the lead exchange for signal generation
            if exchange == self.lead_exchange:
                stats = self._get_btc_stats()
                ret = stats.update(price, btc_return=None)

                # Track BTC log price for spread calculation
                self._latest_btc_log_price = math.log(price) if price > 0 else 0.0

                if ret is not None:
                    self._latest_btc_return = ret
                    self._latest_btc_z = stats.z_score(ret)
                    self._btc_return_count += 1  # Mark as new BTC return

                    # Check if BTC made a significant move
                    if abs(self._latest_btc_z) >= self.config.leader_z_threshold:
                        self._btc_signal_time_ms = now_ms
                        self.logger.debug(
                            f"BTC signal @ {exchange}: Z={self._latest_btc_z:.2f}, "
                            f"return={ret*100:.3f}%"
                        )
        else:
            # Track altcoins per-exchange for cross-exchange arbitrage
            stats = self._get_lagger_stats(exchange, symbol)

            # Only pass BTC return if it's new for this (exchange, symbol) pair.
            # This prevents the same BTC return being used multiple times,
            # which would cause zero BTC variance in the altcoin's stats.
            key = (exchange, symbol)
            last_used = self._last_btc_count_used.get(key, -1)
            if self._btc_return_count > last_used and self._latest_btc_return is not None:
                btc_return = self._latest_btc_return
                self._last_btc_count_used[key] = self._btc_return_count
            else:
                btc_return = None  # Don't update covariance with stale BTC return

            # Pass BTC log price for spread (OU) calculation
            ret = stats.update(
                price,
                btc_return=btc_return,
                btc_log_price=self._latest_btc_log_price,
            )

            if ret is not None and stats.is_ready():
                await self._check_lead_lag_signal(tick, stats, ret, now_ms, price)

    async def _check_lead_lag_signal(
        self,
        tick: Tick,
        stats: RollingStats,
        current_return: float,
        now_ms: int,
        current_price: float,
    ) -> None:
        """
        Check if conditions are met for a lead-lag signal on this specific exchange.

        Signal conditions:
        1. BTC made a significant move recently (within lag window) on lead exchange
        2. This exchange's altcoin hasn't moved proportionally yet
        3. Correlation is strong enough to trust the relationship

        This is the key fix: we check each exchange's altcoin independently,
        so we can catch opportunities where Coinbase ETH is lagging even
        though Binance ETH has already moved.
        """
        symbol = tick.symbol
        exchange = tick.exchange

        # Check if BTC signal is still valid (within lag window)
        time_since_btc_signal = now_ms - self._btc_signal_time_ms
        if time_since_btc_signal > self.config.max_lag_ms:
            return

        if self._latest_btc_return is None:
            return

        # Rate limit signals per (exchange, symbol) pair
        signal_key = (exchange, symbol)
        last_signal = self._last_signal_time.get(signal_key, 0)
        if now_ms - last_signal < 5000:  # 5 second cooldown
            return

        # Get beta and correlation
        beta = stats.beta
        correlation = stats.correlation

        # Bug Fix #3: HARD STOP when correlation is weak
        # If correlation is below threshold, the lead-lag relationship doesn't exist.
        # Do NOT proceed - the model would be gambling on random noise.
        if abs(correlation) < self.config.min_correlation:
            return  # Hard stop - no signal

        # Bug Fix #2: Verify beta and correlation have same sign
        # Mathematically: β = ρ × (σ_alt / σ_btc), so signs must match
        # If they don't, the statistics are unreliable (insufficient data)
        if (beta > 0) != (correlation > 0):
            self.logger.debug(
                f"Beta/correlation sign mismatch for {symbol}@{exchange}: "
                f"β={beta:.4f}, ρ={correlation:.4f} - skipping"
            )
            return  # Hard stop - statistics are broken

        correlation_factor = 1.0  # No reduction needed since we hard-stop on weak correlation

        # Calculate expected return based on BTC move
        expected_return = beta * self._latest_btc_return

        # Calculate return gap (how much altcoin is "behind")
        # This is the residual: ε = R_actual - β × R_btc
        return_gap = expected_return - current_return

        # CRITICAL FIX: Normalize gap by IDIOSYNCRATIC volatility, not total
        #
        # The residual ε should be compared against σ_ε (idiosyncratic vol),
        # not σ_total. Total vol includes systematic risk from BTC.
        #
        # σ_ε = σ_total × √(1 - ρ²)
        #
        # When ρ=0.9: σ_ε ≈ 0.43 × σ_total
        # Using σ_total would understate Z by 2.3x, missing 60%+ of trades.
        gap_z = stats.residual_z_score(return_gap)

        # Check if gap is significant
        if abs(gap_z) < self.config.gap_threshold:
            return

        # Check if altcoin hasn't already moved too much
        altcoin_z = stats.z_score(current_return)
        if abs(altcoin_z) > self.config.lag_z_threshold:
            return  # Altcoin already moved, opportunity missed

        # Bug Fix #4: Fee filter - don't take trades that can't pay for execution
        # Expected return must exceed minimum threshold (2x round-trip fees)
        if abs(expected_return * 100) < self.config.min_expected_return_pct:
            return  # Expected return too small to cover fees

        # Calculate signal confidence
        confidence = self._calculate_confidence(
            leader_z=abs(self._latest_btc_z),
            correlation=abs(correlation),
            gap_z=abs(gap_z),
            time_since_signal_ms=time_since_btc_signal,
            correlation_factor=correlation_factor,
        )

        if confidence < self.config.min_confidence:
            return

        # Determine direction
        if return_gap > 0:
            direction = SignalDirection.LONG  # Altcoin should rise
        else:
            direction = SignalDirection.SHORT  # Altcoin should fall

        # Calculate max hold time based on half-life
        half_life_ms = self._estimate_half_life(symbol, stats)
        max_hold_ms = int(half_life_ms * self.config.half_life_multiplier)

        # Create signal with exchange-specific info
        # Bug Fix #1: Include volatility stats for proper logging
        signal = LeadLagSignal(
            timestamp_ms=now_ms,
            leader_symbol="BTC",
            lagger_symbol=symbol,
            direction=direction,
            leader_return_pct=self._latest_btc_return * 100,
            leader_z_score=self._latest_btc_z,
            lagger_return_pct=current_return * 100,
            expected_return_pct=expected_return * 100,
            return_gap_pct=return_gap * 100,
            rolling_beta=beta,
            rolling_correlation=correlation,
            confidence=confidence,
            max_hold_ms=max_hold_ms,
            exchange=exchange,
            entry_price=current_price,
            total_std=stats.std,                    # σ_total for logging
            idiosyncratic_std=stats.idiosyncratic_std,  # σ_ε for logging
        )

        self._signals_emitted += 1
        self._last_signal_time[signal_key] = now_ms

        self.logger.info(
            f"LEAD-LAG SIGNAL #{self._signals_emitted} @ {exchange}: {signal}"
        )

        if self.on_signal:
            await self.on_signal(signal)

    def _calculate_confidence(
        self,
        leader_z: float,
        correlation: float,
        gap_z: float,
        time_since_signal_ms: int,
        correlation_factor: float,
    ) -> float:
        """
        Calculate signal confidence score (0-1).

        Factors:
        - Strength of BTC move (leader_z)
        - Correlation reliability
        - Size of return gap
        - Time decay (signal weakens over time)
        """
        # Base confidence from leader Z-score (diminishing returns above 3)
        z_conf = min(leader_z / 3.0, 1.0)

        # Correlation factor (linear)
        corr_conf = min(correlation / 0.8, 1.0)

        # Gap factor (bigger gap = more confident)
        gap_conf = min(gap_z / 3.0, 1.0)

        # Time decay (exponential decay)
        time_s = time_since_signal_ms / 1000
        time_conf = math.exp(-self.config.confidence_decay_per_sec * time_s)

        # Combine factors (geometric mean)
        raw_confidence = (z_conf * corr_conf * gap_conf * time_conf) ** 0.25

        # Apply correlation factor adjustment
        return raw_confidence * correlation_factor

    def _estimate_half_life(self, symbol: str, stats: EMAStats) -> float:
        """
        Estimate half-life of spread mean reversion using Ornstein-Uhlenbeck model.

        CRITICAL: We run AR(1) on the SPREAD (cointegration residual), NOT returns.

        The spread S_t = ln(P_alt) - β × ln(P_btc) follows an OU process:
            dS = θ(μ - S)dt + σdW

        In discrete form: S_t = λ × S_{t-1} + (1-λ)μ + ε
        where λ = e^(-θΔt)

        Half-life = ln(2) / θ = -ln(2) / ln(λ)

        Running AR(1) on returns gives the half-life of the BID-ASK BOUNCE,
        which is meaningless for stat arb. The spread is what mean-reverts.

        Returns half-life in milliseconds.
        """
        # Check cache first
        cache_key = f"{symbol}_{id(stats)}"
        if cache_key in self._half_lives:
            return self._half_lives[cache_key]

        # Use spread history (cointegration residual)
        spreads = list(stats.spreads)

        if len(spreads) < 30:
            # Not enough spread data, use default
            return float(self.config.default_half_life_ms)

        # AR(1) regression on the SPREAD: S_t = λ × S_{t-1} + c + ε
        # λ ≈ Cov(S_t, S_{t-1}) / Var(S_{t-1})
        n = len(spreads)
        s_t = spreads[1:]      # S_t
        s_t1 = spreads[:-1]    # S_{t-1}

        mean_t = sum(s_t) / len(s_t)
        mean_t1 = sum(s_t1) / len(s_t1)

        # Covariance and variance
        cov = sum((s_t[i] - mean_t) * (s_t1[i] - mean_t1) for i in range(len(s_t))) / (n - 2)
        var = sum((s - mean_t1) ** 2 for s in s_t1) / (n - 2)

        if var < 1e-16:
            return float(self.config.default_half_life_ms)

        lam = cov / var

        # For mean reversion, we need 0 < λ < 1
        # λ ≈ 1 means no mean reversion (random walk)
        # λ < 0 means oscillation (unlikely for spreads)
        if lam >= 0.999 or lam <= 0:
            # No mean reversion detected, use default
            return float(self.config.default_half_life_ms)

        # Half-life in "ticks"
        # t_0.5 = -ln(2) / ln(λ)
        try:
            half_life_ticks = -math.log(2) / math.log(lam)
        except (ValueError, ZeroDivisionError):
            return float(self.config.default_half_life_ms)

        # Convert to milliseconds
        # Assuming ~100ms average tick interval (adjustable)
        tick_interval_ms = 100
        half_life_ms = half_life_ticks * tick_interval_ms

        # Clamp to reasonable range (5s to 120s)
        half_life_ms = max(5000, min(120000, half_life_ms))

        # Cache result (per stats instance to handle multi-exchange)
        self._half_lives[cache_key] = half_life_ms

        self.logger.debug(
            f"Half-life for {symbol}: λ={lam:.4f}, "
            f"t_0.5={half_life_ms/1000:.1f}s ({len(spreads)} samples)"
        )

        return half_life_ms

    def get_stats(self) -> dict:
        """Return current model statistics."""
        # Get unique symbols being tracked
        symbols_tracked = set()
        exchanges_tracked = set()
        for (exchange, symbol) in self._lagger_stats.keys():
            symbols_tracked.add(symbol)
            exchanges_tracked.add(exchange)

        stats = {
            "signals_emitted": self._signals_emitted,
            "symbols_tracked": sorted(symbols_tracked),
            "exchanges_tracked": sorted(exchanges_tracked),
            "latest_btc_z": self._latest_btc_z,
        }

        # Add per (exchange, symbol) stats
        for (exchange, symbol), s in self._lagger_stats.items():
            if s.is_ready():
                key = f"{exchange}_{symbol}"
                stats[f"{key}_beta"] = round(s.beta, 3)
                stats[f"{key}_correlation"] = round(s.correlation, 3)

        return stats
