"""
Model Manager - Orchestrates all quantitative models.

Provides:
- Centralized model instantiation and lifecycle
- Tick distribution to all models
- Aggregated stats collection for dashboard
- Signal routing
"""
import asyncio
import math
import time
from dataclasses import dataclass, field
from typing import Any, Callable, Awaitable

from models import Tick


def sanitize_for_json(obj: Any) -> Any:
    """Replace inf/nan values with None for JSON serialization."""
    if isinstance(obj, dict):
        return {k: sanitize_for_json(v) for k, v in obj.items()}
    elif isinstance(obj, list):
        return [sanitize_for_json(item) for item in obj]
    elif isinstance(obj, float):
        if math.isinf(obj) or math.isnan(obj):
            return None
        return obj
    return obj
from .base import BaseModel, ModelSignal, SignalType
from .volatility.garch import GARCH11, EGARCH
from .volatility.realized_vol import RealizedVolatility
from .risk.var_es import RealTimeVaR, EVTVaR
from .regime.hamilton_filter import HamiltonFilter, MSGARCHFilter
from .jumps.jump_detection import LeeMyklandTest
from .stat_arb.kalman_filter import KalmanHedgeRatio
from .microstructure.price_impact import KyleLambda


@dataclass
class ModelConfig:
    """Configuration for model instantiation."""
    symbols: list[str] = field(default_factory=lambda: ["BTC", "ETH", "SOL"])
    enable_garch: bool = True
    enable_egarch: bool = True
    enable_realized_vol: bool = True
    enable_var: bool = True
    enable_evt: bool = True
    enable_regime: bool = True
    enable_msgarch: bool = False  # Heavy, optional
    enable_jumps: bool = True
    enable_kalman: bool = True
    enable_microstructure: bool = True
    warmup_ticks: int = 200


class ModelManager:
    """
    Central manager for all quantitative models.

    Usage:
        manager = ModelManager(config)
        await manager.initialize()

        # Route ticks from exchange connectors
        await manager.on_tick(tick)

        # Get stats for dashboard
        stats = manager.get_all_stats()
    """

    def __init__(
        self,
        config: ModelConfig | None = None,
        on_signal: Callable[[ModelSignal], Awaitable[None]] | None = None,
    ):
        self.config = config or ModelConfig()
        self.on_signal = on_signal

        # Model storage by category
        self.volatility_models: dict[str, BaseModel] = {}
        self.risk_models: dict[str, BaseModel] = {}
        self.regime_models: dict[str, BaseModel] = {}
        self.jump_models: dict[str, BaseModel] = {}
        self.stat_arb_models: dict[str, BaseModel] = {}
        self.microstructure_models: dict[str, BaseModel] = {}

        # Combined for iteration
        self._all_models: list[BaseModel] = []

        # Stats
        self._tick_count = 0
        self._start_time = time.time()
        self._signal_count = 0
        self._recent_signals: list[dict] = []  # Last 50 signals

    async def initialize(self) -> None:
        """Initialize all enabled models."""
        cfg = self.config

        for symbol in cfg.symbols:
            # Volatility models
            if cfg.enable_garch:
                model = GARCH11(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.volatility_models[f"garch_{symbol}"] = model
                self._all_models.append(model)

            if cfg.enable_egarch:
                model = EGARCH(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.volatility_models[f"egarch_{symbol}"] = model
                self._all_models.append(model)

            if cfg.enable_realized_vol:
                model = RealizedVolatility(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.volatility_models[f"rv_{symbol}"] = model
                self._all_models.append(model)

            # Risk models
            if cfg.enable_var:
                # Link GARCH to VaR if available
                garch = self.volatility_models.get(f"garch_{symbol}")
                model = RealTimeVaR(
                    symbol=symbol,
                    garch_model=garch,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.risk_models[f"var_{symbol}"] = model
                self._all_models.append(model)

            if cfg.enable_evt:
                model = EVTVaR(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks * 2,  # EVT needs more data
                    on_signal=self._handle_signal,
                )
                self.risk_models[f"evt_{symbol}"] = model
                self._all_models.append(model)

            # Regime models
            if cfg.enable_regime:
                model = HamiltonFilter(
                    symbol=symbol,
                    n_regimes=2,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.regime_models[f"regime_{symbol}"] = model
                self._all_models.append(model)

            if cfg.enable_msgarch:
                model = MSGARCHFilter(
                    symbol=symbol,
                    n_regimes=2,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.regime_models[f"msgarch_{symbol}"] = model
                self._all_models.append(model)

            # Jump detection
            if cfg.enable_jumps:
                model = LeeMyklandTest(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.jump_models[f"jumps_{symbol}"] = model
                self._all_models.append(model)

            # Microstructure
            if cfg.enable_microstructure:
                model = KyleLambda(
                    symbol=symbol,
                    warmup_ticks=cfg.warmup_ticks,
                    on_signal=self._handle_signal,
                )
                self.microstructure_models[f"kyle_{symbol}"] = model
                self._all_models.append(model)

        # Pair-based models (Kalman for hedge ratios)
        if cfg.enable_kalman and len(cfg.symbols) >= 2:
            # Create ETH-BTC and SOL-BTC pairs
            base = "BTC"
            for symbol in cfg.symbols:
                if symbol != base:
                    model = KalmanHedgeRatio(
                        symbol=symbol,
                        warmup_ticks=cfg.warmup_ticks,
                        on_signal=self._handle_signal,
                    )
                    self.stat_arb_models[f"kalman_{symbol}_{base}"] = model
                    self._all_models.append(model)

    async def _handle_signal(self, signal: ModelSignal) -> None:
        """Handle signal from any model."""
        self._signal_count += 1

        # Store recent signals
        signal_dict = {
            "timestamp": signal.timestamp_ms,
            "model": signal.model_name,
            "type": signal.signal_type.value,
            "symbol": signal.symbol,
            "confidence": signal.confidence,
            "metadata": signal.metadata,
        }
        self._recent_signals.append(signal_dict)
        if len(self._recent_signals) > 50:
            self._recent_signals = self._recent_signals[-50:]

        # Forward to external handler
        if self.on_signal:
            await self.on_signal(signal)

    async def on_tick(self, tick: Tick) -> None:
        """
        Route tick to all relevant models.

        This is called for every tick from the exchange connectors.
        """
        self._tick_count += 1

        # Send to all models - they filter by symbol internally
        tasks = [model.on_tick(tick) for model in self._all_models]
        await asyncio.gather(*tasks, return_exceptions=True)

    def get_all_stats(self) -> dict[str, Any]:
        """
        Get aggregated stats from all models for dashboard.

        Returns nested dict organized by category.
        """
        stats = {
            "manager": {
                "tick_count": self._tick_count,
                "runtime_seconds": time.time() - self._start_time,
                "total_models": len(self._all_models),
                "signal_count": self._signal_count,
            },
            "volatility": {
                name: model.get_stats()
                for name, model in self.volatility_models.items()
            },
            "risk": {
                name: model.get_stats()
                for name, model in self.risk_models.items()
            },
            "regime": {
                name: model.get_stats()
                for name, model in self.regime_models.items()
            },
            "jumps": {
                name: model.get_stats()
                for name, model in self.jump_models.items()
            },
            "stat_arb": {
                name: model.get_stats()
                for name, model in self.stat_arb_models.items()
            },
            "microstructure": {
                name: model.get_stats()
                for name, model in self.microstructure_models.items()
            },
            "recent_signals": self._recent_signals[-20:],
        }
        # Sanitize inf/nan values for JSON serialization
        return sanitize_for_json(stats)

    def get_model_stats(self, model_name: str) -> dict | None:
        """Get stats for a specific model by name."""
        for model in self._all_models:
            if model.name == model_name:
                return sanitize_for_json(model.get_stats())
        return None

    def get_category_stats(self, category: str) -> dict:
        """Get stats for a model category."""
        category_map = {
            "volatility": self.volatility_models,
            "risk": self.risk_models,
            "regime": self.regime_models,
            "jumps": self.jump_models,
            "stat_arb": self.stat_arb_models,
            "microstructure": self.microstructure_models,
        }
        models = category_map.get(category, {})
        stats = {name: model.get_stats() for name, model in models.items()}
        return sanitize_for_json(stats)

    def get_symbol_stats(self, symbol: str) -> dict:
        """Get all model stats for a specific symbol."""
        stats = {}
        for model in self._all_models:
            model_stats = model.get_stats()
            if model_stats.get("symbol") == symbol:
                stats[model.name] = model_stats
        return sanitize_for_json(stats)
