"""
Base model interface for quantitative financial models.

All models follow the LeadLagBrain pattern:
- O(1) on_tick() updates (< 1ms per call)
- EMA-based statistics for numerical stability
- Warmup period before signal generation
- get_stats() for dashboard transparency

Mathematical foundations from the treatise are implemented with
crypto-specific adaptations for high volatility and 24/7 markets.
"""
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Callable, Awaitable, Any
from enum import Enum
import time

from models import Tick


class SignalType(Enum):
    """Types of signals emitted by quantitative models."""
    VOL_SPIKE = "vol_spike"           # Volatility spike detected
    VOL_REGIME_CHANGE = "vol_regime"  # Volatility regime shift
    REGIME_CHANGE = "regime_change"   # Market regime transition
    JUMP_DETECTED = "jump_detected"   # Price jump detected
    RISK_BREACH = "risk_breach"       # VaR/ES threshold breached
    CORRELATION_BREAK = "corr_break"  # Correlation breakdown


@dataclass(slots=True)
class ModelSignal:
    """
    Signal emitted by any quantitative model.

    Follows the LeadLagSignal pattern for consistency
    with the existing arbitrage system.
    """
    timestamp_ms: int
    model_name: str
    signal_type: SignalType
    symbol: str
    confidence: float  # 0.0 to 1.0

    # Model-specific data
    metadata: dict = field(default_factory=dict)

    # Optional: direction for tradeable signals
    direction: str = ""  # "long", "short", or ""

    def __repr__(self) -> str:
        return (
            f"ModelSignal({self.model_name}: {self.signal_type.value} "
            f"on {self.symbol}, conf={self.confidence:.2f})"
        )


# Type alias for signal callbacks
SignalCallback = Callable[[ModelSignal], Awaitable[None]]


class BaseModel(ABC):
    """
    Abstract base class for all quantitative models.

    Design principles (from existing LeadLagBrain):
    1. O(1) per-tick updates - no O(n) operations in on_tick()
    2. EMA-based statistics for numerical stability
    3. Warmup period before generating signals
    4. get_stats() returns JSON-serializable dict for dashboard

    Subclasses must implement:
    - on_tick(): Process incoming market tick
    - get_stats(): Return current model state

    Optional overrides:
    - _emit_signal(): Customize signal emission
    - reset(): Reset model state
    """

    # Use __slots__ in subclasses for memory efficiency

    def __init__(
        self,
        name: str,
        warmup_ticks: int = 100,
        on_signal: SignalCallback | None = None,
    ):
        """
        Initialize base model.

        Args:
            name: Unique model identifier (e.g., "garch_BTC")
            warmup_ticks: Minimum ticks before emitting signals
            on_signal: Async callback for signal emission
        """
        self.name = name
        self.warmup_ticks = warmup_ticks
        self.on_signal = on_signal

        # State tracking
        self._tick_count: int = 0
        self._is_ready: bool = False
        self._signals_emitted: int = 0
        self._last_signal_time_ms: int = 0

        # Rate limiting: minimum ms between signals
        self._signal_cooldown_ms: int = 1000

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

        MUST complete in < 1ms for real-time performance.
        Update model state and emit signals if conditions met.

        Args:
            tick: Normalized tick from exchange connector
        """
        pass

    @abstractmethod
    def get_stats(self) -> dict:
        """
        Return current model state for dashboard display.

        Must return JSON-serializable dict with at minimum:
        - "model": model name
        - "is_ready": bool
        - "tick_count": int

        Add model-specific statistics as needed.
        """
        pass

    def is_ready(self) -> bool:
        """Check if warmup complete and model can emit signals."""
        return self._is_ready

    def reset(self) -> None:
        """Reset model state. Override in subclasses for full reset."""
        self._tick_count = 0
        self._is_ready = False
        self._signals_emitted = 0
        self._last_signal_time_ms = 0

    async def _emit_signal(
        self,
        signal_type: SignalType,
        symbol: str,
        confidence: float,
        metadata: dict | None = None,
        direction: str = "",
    ) -> None:
        """
        Emit a model signal via callback.

        Handles rate limiting and signal counting.

        Args:
            signal_type: Type of signal being emitted
            symbol: Asset symbol
            confidence: Signal confidence (0-1)
            metadata: Model-specific data
            direction: Optional trade direction
        """
        if not self.on_signal:
            return

        now_ms = int(time.time() * 1000)

        # Rate limiting
        if now_ms - self._last_signal_time_ms < self._signal_cooldown_ms:
            return

        signal = ModelSignal(
            timestamp_ms=now_ms,
            model_name=self.name,
            signal_type=signal_type,
            symbol=symbol,
            confidence=min(max(confidence, 0.0), 1.0),
            metadata=metadata or {},
            direction=direction,
        )

        self._signals_emitted += 1
        self._last_signal_time_ms = now_ms

        await self.on_signal(signal)

    def _check_warmup(self) -> None:
        """Update warmup status after processing tick."""
        self._tick_count += 1
        if self._tick_count >= self.warmup_ticks and not self._is_ready:
            self._is_ready = True

    def _base_stats(self) -> dict:
        """Return base statistics common to all models."""
        return {
            "model": self.name,
            "is_ready": self._is_ready,
            "tick_count": self._tick_count,
            "signals_emitted": self._signals_emitted,
        }


class MultiSymbolModel(BaseModel):
    """
    Base class for models that track multiple symbols.

    Extends BaseModel with per-symbol state management.
    """

    def __init__(
        self,
        name: str,
        symbols: list[str],
        warmup_ticks: int = 100,
        on_signal: SignalCallback | None = None,
    ):
        super().__init__(name, warmup_ticks, on_signal)
        self.symbols = symbols
        self._symbol_tick_counts: dict[str, int] = {s: 0 for s in symbols}
        self._symbol_ready: dict[str, bool] = {s: False for s in symbols}

    def is_symbol_ready(self, symbol: str) -> bool:
        """Check if specific symbol has completed warmup."""
        return self._symbol_ready.get(symbol, False)

    def _check_symbol_warmup(self, symbol: str) -> None:
        """Update warmup status for specific symbol."""
        if symbol not in self._symbol_tick_counts:
            self._symbol_tick_counts[symbol] = 0
            self._symbol_ready[symbol] = False

        self._symbol_tick_counts[symbol] += 1
        self._tick_count += 1

        if self._symbol_tick_counts[symbol] >= self.warmup_ticks:
            self._symbol_ready[symbol] = True

        # Overall model ready when all symbols ready
        if all(self._symbol_ready.values()) and not self._is_ready:
            self._is_ready = True
