"""
Real-time Web Dashboard Server.

Provides WebSocket streaming of arbitrage data and REST API for historical data.
Uses FastAPI for high-performance async operations.
"""
import asyncio
import json
import time
from collections import deque
from contextlib import asynccontextmanager
from dataclasses import dataclass, field, asdict
from datetime import datetime
from decimal import Decimal
from enum import Enum
from pathlib import Path
from typing import Any

from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import HTMLResponse, FileResponse
from fastapi.staticfiles import StaticFiles

from config import SYMBOLS, LEAD_EXCHANGE, LAG_EXCHANGES, FEES, SPREAD_THRESHOLD_BPS
from models import Tick
from engine.lead_lag_brain import LeadLagSignal, SignalDirection, RollingStats
from engine.spread_calculator import SpreadResult
from utils import get_logger


logger = get_logger("web.server")


# ============================================================================
# Data Models for WebSocket Messages
# ============================================================================

class MessageType(str, Enum):
    """WebSocket message types."""
    TICK = "tick"
    SPREAD = "spread"
    LEAD_LAG_SIGNAL = "lead_lag_signal"
    LEAD_LAG_STATE = "lead_lag_state"
    STATS = "stats"
    TRADE = "trade"
    ERROR = "error"
    CONNECTED = "connected"
    PORTFOLIO_STATE = "portfolio_state"
    TRADE_SIZE_UPDATE = "trade_size_update"


@dataclass
class TickMessage:
    """Tick data for WebSocket."""
    type: str = MessageType.TICK
    exchange: str = ""
    symbol: str = ""
    bid: float = 0.0
    ask: float = 0.0
    mid: float = 0.0
    spread_bps: float = 0.0
    timestamp: int = 0

    @classmethod
    def from_tick(cls, tick: Tick) -> "TickMessage":
        mid = float(tick.mid_price())
        return cls(
            exchange=tick.exchange,
            symbol=tick.symbol,
            bid=float(tick.bid),
            ask=float(tick.ask),
            mid=mid,
            spread_bps=float(tick.spread_bps()),
            timestamp=tick.local_ts,
        )


@dataclass
class SpreadMessage:
    """Spread calculation for WebSocket."""
    type: str = MessageType.SPREAD
    symbol: str = ""
    buy_exchange: str = ""
    sell_exchange: str = ""
    buy_price: float = 0.0
    sell_price: float = 0.0
    gross_spread_bps: float = 0.0
    net_spread_bps: float = 0.0
    buy_fee_pct: float = 0.0
    sell_fee_pct: float = 0.0
    is_opportunity: bool = False
    timestamp: int = 0

    @classmethod
    def from_spread_result(cls, spread: SpreadResult, fees: dict | None = None) -> "SpreadMessage":
        # Use provided fees or fall back to config
        if fees is None:
            fees = FEES
        return cls(
            symbol=spread.symbol,
            buy_exchange=spread.buy_exchange,
            sell_exchange=spread.sell_exchange,
            buy_price=float(spread.buy_price),
            sell_price=float(spread.sell_price),
            gross_spread_bps=float(spread.gross_spread_bps),
            net_spread_bps=float(spread.spread_bps),
            buy_fee_pct=float(fees.get(spread.buy_exchange, Decimal("0.001"))) * 100,
            sell_fee_pct=float(fees.get(spread.sell_exchange, Decimal("0.001"))) * 100,
            is_opportunity=spread.is_opportunity,
            timestamp=int(time.time() * 1000),
        )


@dataclass
class LeadLagStateMessage:
    """
    Complete state of the lead-lag model for transparency.

    Shows all intermediate calculations so users can verify the math.
    """
    type: str = MessageType.LEAD_LAG_STATE
    timestamp: int = 0

    # BTC (Leader) State
    btc_price: float = 0.0
    btc_return: float = 0.0
    btc_return_mean: float = 0.0
    btc_return_std: float = 0.0
    btc_z_score: float = 0.0
    btc_is_signal: bool = False

    # Per-symbol state
    symbols: dict = field(default_factory=dict)

    # Thresholds (for display)
    leader_z_threshold: float = 2.0
    lag_z_threshold: float = 1.0
    gap_threshold: float = 1.5
    min_correlation: float = 0.5


@dataclass
class SymbolState:
    """State for a single altcoin symbol."""
    symbol: str = ""
    price: float = 0.0
    return_pct: float = 0.0
    return_mean: float = 0.0
    return_std: float = 0.0
    z_score: float = 0.0

    # Beta calculation components
    rolling_beta: float = 0.0
    rolling_correlation: float = 0.0
    covariance: float = 0.0
    btc_variance: float = 0.0

    # Expected vs actual
    expected_return: float = 0.0
    return_gap: float = 0.0
    gap_z_score: float = 0.0

    # Signal state
    has_signal: bool = False
    signal_direction: str = ""
    signal_confidence: float = 0.0


@dataclass
class LeadLagSignalMessage:
    """Lead-lag signal for WebSocket."""
    type: str = MessageType.LEAD_LAG_SIGNAL
    timestamp: int = 0
    symbol: str = ""
    direction: str = ""
    btc_return_pct: float = 0.0
    btc_z_score: float = 0.0
    expected_return_pct: float = 0.0
    actual_return_pct: float = 0.0
    return_gap_pct: float = 0.0
    rolling_beta: float = 0.0
    rolling_correlation: float = 0.0
    confidence: float = 0.0
    max_hold_ms: int = 0

    @classmethod
    def from_signal(cls, signal: LeadLagSignal) -> "LeadLagSignalMessage":
        return cls(
            timestamp=signal.timestamp_ms,
            symbol=signal.lagger_symbol,
            direction=signal.direction.value,
            btc_return_pct=signal.leader_return_pct,
            btc_z_score=signal.leader_z_score,
            expected_return_pct=signal.expected_return_pct,
            actual_return_pct=signal.lagger_return_pct,
            return_gap_pct=signal.return_gap_pct,
            rolling_beta=signal.rolling_beta,
            rolling_correlation=signal.rolling_correlation,
            confidence=signal.confidence,
            max_hold_ms=signal.max_hold_ms,
        )


@dataclass
class StatsMessage:
    """Aggregated statistics for WebSocket."""
    type: str = MessageType.STATS
    timestamp: int = 0
    runtime_seconds: float = 0.0
    ticks_received: dict = field(default_factory=dict)
    spread_opportunities: int = 0
    lead_lag_signals: int = 0
    spread_pnl_usd: float = 0.0
    lead_lag_pnl_usd: float = 0.0
    total_pnl_usd: float = 0.0


@dataclass
class PortfolioStateMessage:
    """Portfolio state for dashboard display."""
    type: str = MessageType.PORTFOLIO_STATE
    timestamp: int = 0
    trade_size_usd: float = 0.0
    currency: str = "USD"
    spread_pnl_usd: float = 0.0
    lead_lag_pnl_usd: float = 0.0
    total_pnl_usd: float = 0.0
    pending_positions: int = 0
    portfolio_value_usd: float = 0.0


# GBP to USD conversion rate - updated via live API
GBP_TO_USD_RATE = 1.27  # Fallback default

async def fetch_live_exchange_rate() -> float:
    """Fetch live GBP/USD exchange rate from API."""
    import aiohttp

    try:
        async with aiohttp.ClientSession() as session:
            # Try primary API (exchangerate.host)
            async with session.get(
                'https://api.exchangerate.host/latest?base=GBP&symbols=USD',
                timeout=aiohttp.ClientTimeout(total=5)
            ) as response:
                if response.status == 200:
                    data = await response.json()
                    if data.get('success') and data.get('rates', {}).get('USD'):
                        return data['rates']['USD']
    except Exception as e:
        logger.warning(f"Primary exchange rate API failed: {e}")

    try:
        async with aiohttp.ClientSession() as session:
            # Try alternative API (open.er-api.com)
            async with session.get(
                'https://open.er-api.com/v6/latest/GBP',
                timeout=aiohttp.ClientTimeout(total=5)
            ) as response:
                if response.status == 200:
                    data = await response.json()
                    if data.get('result') == 'success' and data.get('rates', {}).get('USD'):
                        return data['rates']['USD']
    except Exception as e:
        logger.warning(f"Alternative exchange rate API failed: {e}")

    return GBP_TO_USD_RATE  # Return fallback


# ============================================================================
# WebSocket Connection Manager
# ============================================================================

class ConnectionManager:
    """Manages WebSocket connections and broadcasts."""

    def __init__(self):
        self.active_connections: list[WebSocket] = []
        self._lock = asyncio.Lock()

    async def connect(self, websocket: WebSocket):
        await websocket.accept()
        async with self._lock:
            self.active_connections.append(websocket)
        logger.info(f"Client connected. Total: {len(self.active_connections)}")

    async def disconnect(self, websocket: WebSocket):
        async with self._lock:
            if websocket in self.active_connections:
                self.active_connections.remove(websocket)
        logger.info(f"Client disconnected. Total: {len(self.active_connections)}")

    async def broadcast(self, message: dict):
        """Broadcast message to all connected clients."""
        if not self.active_connections:
            return

        data = json.dumps(message, default=str)
        async with self._lock:
            dead_connections = []
            for connection in self.active_connections:
                try:
                    await connection.send_text(data)
                except Exception:
                    dead_connections.append(connection)

            for conn in dead_connections:
                self.active_connections.remove(conn)


# ============================================================================
# Dashboard State
# ============================================================================

class DashboardState:
    """
    Holds all state needed for the dashboard.

    This is the bridge between the arbitrage engine and the web frontend.
    """

    def __init__(self):
        self.start_time = time.time()
        self.manager = ConnectionManager()

        # Latest ticks per exchange/symbol
        self.latest_ticks: dict[str, dict[str, TickMessage]] = {}

        # Latest spreads
        self.latest_spreads: dict[str, SpreadMessage] = {}  # key: "symbol:buy->sell"

        # Lead-lag model state (for transparency)
        self.lead_lag_state = LeadLagStateMessage()

        # History buffers (for charts)
        self.price_history: dict[str, deque] = {}  # symbol -> [(ts, price), ...]
        self.spread_history: deque = deque(maxlen=500)
        self.signal_history: deque = deque(maxlen=100)

        # Stats
        self.ticks_received: dict[str, int] = {}
        self.spread_opportunities = 0
        self.lead_lag_signals = 0
        self.spread_pnl = 0.0
        self.lead_lag_pnl = 0.0

        # Portfolio / Paper trading state
        self.trade_size_usd = 1000.0  # Default trade size
        self.trade_currency = "USD"   # User's selected currency
        self._paper_trader_ref = None  # Reference to paper trader for dynamic updates

        # Dynamic fees and threshold (adjustable from dashboard)
        self.fees: dict[str, Decimal] = dict(FEES)  # Copy of config fees
        self.spread_threshold_bps: Decimal = SPREAD_THRESHOLD_BPS
        self.lead_lag_min_return_pct: float = 0.25  # Lead-lag model minimum expected return (%)
        self._detector_ref = None  # Reference to arbitrage detector for dynamic updates

        # Quantitative models manager
        self._model_manager = None  # Reference to ModelManager for quant models

    def set_model_manager(self, model_manager) -> None:
        """Set reference to model manager for quantitative models."""
        self._model_manager = model_manager

    def set_paper_trader(self, paper_trader) -> None:
        """Set reference to paper trader for dynamic updates."""
        self._paper_trader_ref = paper_trader
        self.trade_size_usd = float(paper_trader.trade_size_usd)

    def set_detector(self, detector) -> None:
        """Set reference to arbitrage detector for dynamic fee/threshold updates."""
        self._detector_ref = detector
        self.spread_threshold_bps = detector.threshold_bps
        self.fees = detector.fees  # Sync fees from detector
        self.lead_lag_min_return_pct = detector.lead_lag_min_return_pct  # Sync lead-lag threshold

    async def update_fees(self, exchange: str, fee_pct: float) -> dict:
        """
        Update fee for a specific exchange.

        Args:
            exchange: Exchange name (binance, coinbase, bybit)
            fee_pct: Fee percentage (e.g., 0.1 for 0.1%)

        Returns:
            Dict with update result
        """
        if exchange not in self.fees:
            return {"error": f"Unknown exchange: {exchange}", "status": "error"}

        if fee_pct < 0 or fee_pct > 5:
            return {"error": "Fee must be between 0% and 5%", "status": "error"}

        old_fee = float(self.fees[exchange]) * 100
        new_fee_decimal = Decimal(str(fee_pct / 100))  # Convert % to decimal
        self.fees[exchange] = new_fee_decimal

        # Update detector if connected
        if self._detector_ref:
            self._detector_ref.fees[exchange] = new_fee_decimal

        logger.info(f"Fee updated for {exchange}: {old_fee:.3f}% -> {fee_pct:.3f}%")

        # Broadcast update to all clients
        await self.broadcast_settings_update()

        return {
            "status": "success",
            "exchange": exchange,
            "old_fee_pct": old_fee,
            "new_fee_pct": fee_pct,
        }

    async def update_spread_threshold(self, threshold_bps: float) -> dict:
        """
        Update spread threshold in basis points.

        Args:
            threshold_bps: New threshold in basis points

        Returns:
            Dict with update result
        """
        if threshold_bps < 0 or threshold_bps > 1000:
            return {"error": "Threshold must be between 0 and 1000 bps", "status": "error"}

        old_threshold = float(self.spread_threshold_bps)
        self.spread_threshold_bps = Decimal(str(threshold_bps))

        # Update detector if connected
        if self._detector_ref:
            self._detector_ref.threshold_bps = self.spread_threshold_bps

        logger.info(f"Spread threshold updated: {old_threshold:.1f}bps -> {threshold_bps:.1f}bps")

        # Broadcast update to all clients
        await self.broadcast_settings_update()

        return {
            "status": "success",
            "old_threshold_bps": old_threshold,
            "new_threshold_bps": threshold_bps,
        }

    async def update_lead_lag_min_return(self, min_return_pct: float) -> dict:
        """
        Update the lead-lag model's minimum expected return threshold.

        Args:
            min_return_pct: Minimum expected return in percentage (e.g., 0.25 for 0.25%)

        Returns:
            Dict with update result
        """
        if min_return_pct < 0 or min_return_pct > 10:
            return {"error": "Min return must be between 0% and 10%", "status": "error"}

        old_value = self.lead_lag_min_return_pct
        self.lead_lag_min_return_pct = min_return_pct

        # Update detector if connected
        if self._detector_ref:
            self._detector_ref.lead_lag_min_return_pct = min_return_pct

        logger.info(f"Lead-lag min return updated: {old_value:.3f}% -> {min_return_pct:.3f}%")

        # Broadcast update to all clients
        await self.broadcast_settings_update()

        return {
            "status": "success",
            "old_min_return_pct": old_value,
            "new_min_return_pct": min_return_pct,
        }

    async def broadcast_settings_update(self) -> None:
        """Broadcast current settings (fees and threshold) to all clients."""
        msg = {
            "type": "settings_update",
            "timestamp": int(time.time() * 1000),
            "fees": {k: float(v) * 100 for k, v in self.fees.items()},
            "spread_threshold_bps": float(self.spread_threshold_bps),
            "lead_lag_min_return_pct": self.lead_lag_min_return_pct,
        }
        await self.manager.broadcast(msg)

    async def update_trade_size(self, new_size: float, currency: str = "USD") -> dict:
        """
        Update trade size and broadcast to all clients.

        Args:
            new_size: Trade size in the specified currency
            currency: "USD" or "GBP"

        Returns:
            Dict with update result
        """
        global GBP_TO_USD_RATE

        if self._paper_trader_ref is None:
            return {"error": "Paper trader not initialized", "status": "error"}

        # Convert GBP to USD if needed, using live rate
        size_usd = new_size
        if currency == "GBP":
            # Fetch live rate
            live_rate = await fetch_live_exchange_rate()
            GBP_TO_USD_RATE = live_rate
            size_usd = new_size * live_rate
            logger.info(f"Converting £{new_size} to ${size_usd:.2f} (rate: {live_rate:.4f})")

        result = self._paper_trader_ref.set_trade_size(Decimal(str(size_usd)), currency)
        self.trade_size_usd = size_usd
        self.trade_currency = currency

        # Add exchange rate to result
        result["exchange_rate"] = GBP_TO_USD_RATE

        # Broadcast update to all connected clients
        await self.broadcast_portfolio_state()

        result["status"] = "success"
        return result

    async def broadcast_portfolio_state(self) -> None:
        """Broadcast current portfolio state to all connected clients."""
        if self._paper_trader_ref is None:
            return

        state = self._paper_trader_ref.get_portfolio_state()
        total_pnl = state["total_pnl_usd"]
        portfolio_value = self.trade_size_usd + total_pnl

        msg = {
            "type": "portfolio_state",
            "timestamp": int(time.time() * 1000),
            "trade_size_usd": self.trade_size_usd,
            "currency": self.trade_currency,
            "spread_pnl_usd": state["spread_pnl_usd"],
            "lead_lag_pnl_usd": state["lead_lag_pnl_usd"],
            "total_pnl_usd": total_pnl,
            "pending_positions": state["pending_positions"],
            "portfolio_value_usd": portfolio_value,
            "lead_lag_trades_closed": state["lead_lag_trades_closed"],
            "win_rate": state["win_rate"],
        }

        await self.manager.broadcast(msg)

    async def on_tick(self, tick: Tick):
        """Handle incoming tick."""
        msg = TickMessage.from_tick(tick)

        # Store latest
        if tick.exchange not in self.latest_ticks:
            self.latest_ticks[tick.exchange] = {}
        self.latest_ticks[tick.exchange][tick.symbol] = msg

        # Update tick count
        self.ticks_received[tick.exchange] = self.ticks_received.get(tick.exchange, 0) + 1

        # Store price history
        key = f"{tick.exchange}:{tick.symbol}"
        if key not in self.price_history:
            self.price_history[key] = deque(maxlen=500)
        self.price_history[key].append((tick.local_ts, msg.mid))

        # Broadcast tick
        await self.manager.broadcast(asdict(msg))

    async def on_spread(self, spread: SpreadResult):
        """Handle spread calculation."""
        msg = SpreadMessage.from_spread_result(spread, fees=self.fees)

        # Store latest
        key = f"{spread.symbol}:{spread.buy_exchange}->{spread.sell_exchange}"
        self.latest_spreads[key] = msg

        # Store history
        self.spread_history.append(msg)

        if spread.is_opportunity:
            self.spread_opportunities += 1

        # Broadcast
        await self.manager.broadcast(asdict(msg))

    async def on_lead_lag_signal(self, signal: LeadLagSignal):
        """Handle lead-lag signal."""
        msg = LeadLagSignalMessage.from_signal(signal)

        self.lead_lag_signals += 1
        self.signal_history.append(msg)

        # Broadcast
        await self.manager.broadcast(asdict(msg))

    async def update_lead_lag_state(self, state: LeadLagStateMessage):
        """Update complete lead-lag model state."""
        self.lead_lag_state = state
        await self.manager.broadcast(asdict(state))

    async def broadcast_stats(self):
        """Broadcast current statistics."""
        msg = StatsMessage(
            timestamp=int(time.time() * 1000),
            runtime_seconds=time.time() - self.start_time,
            ticks_received=dict(self.ticks_received),
            spread_opportunities=self.spread_opportunities,
            lead_lag_signals=self.lead_lag_signals,
            spread_pnl_usd=self.spread_pnl,
            lead_lag_pnl_usd=self.lead_lag_pnl,
            total_pnl_usd=self.spread_pnl + self.lead_lag_pnl,
        )
        await self.manager.broadcast(asdict(msg))

    def get_initial_state(self) -> dict:
        """Get complete state for new connections."""
        total_pnl = self.spread_pnl + self.lead_lag_pnl
        portfolio_value = self.trade_size_usd + total_pnl

        # Get additional portfolio state if paper trader is connected
        pending_positions = 0
        win_rate = 0.0
        if self._paper_trader_ref:
            state = self._paper_trader_ref.get_portfolio_state()
            pending_positions = state.get("pending_positions", 0)
            win_rate = state.get("win_rate", 0.0)

        return {
            "type": "initial_state",
            "timestamp": int(time.time() * 1000),
            "symbols": SYMBOLS,
            "lead_exchange": LEAD_EXCHANGE,
            "lag_exchanges": LAG_EXCHANGES,
            "fees": {k: float(v) * 100 for k, v in self.fees.items()},
            "spread_threshold_bps": float(self.spread_threshold_bps),
            "lead_lag_min_return_pct": self.lead_lag_min_return_pct,
            "latest_ticks": {
                ex: {sym: asdict(tick) for sym, tick in ticks.items()}
                for ex, ticks in self.latest_ticks.items()
            },
            "latest_spreads": {k: asdict(v) for k, v in self.latest_spreads.items()},
            "lead_lag_state": asdict(self.lead_lag_state),
            "stats": {
                "runtime_seconds": time.time() - self.start_time,
                "ticks_received": dict(self.ticks_received),
                "spread_opportunities": self.spread_opportunities,
                "lead_lag_signals": self.lead_lag_signals,
                "spread_pnl_usd": self.spread_pnl,
                "lead_lag_pnl_usd": self.lead_lag_pnl,
            },
            "price_history": {
                k: list(v)[-100:]  # Last 100 points
                for k, v in self.price_history.items()
            },
            "signal_history": [asdict(s) for s in list(self.signal_history)[-20:]],
            # Portfolio state for paper trading
            "trade_size_usd": self.trade_size_usd,
            "trade_currency": self.trade_currency,
            "portfolio_value_usd": portfolio_value,
            "pending_positions": pending_positions,
            "win_rate": win_rate,
            # Quantitative models state
            "quant_models": self._model_manager.get_all_stats() if self._model_manager else {},
        }

    async def broadcast_model_stats(self) -> None:
        """Broadcast quantitative model statistics to all clients."""
        if self._model_manager is None:
            return

        msg = {
            "type": "model_stats",
            "timestamp": int(time.time() * 1000),
            "models": self._model_manager.get_all_stats(),
        }
        await self.manager.broadcast(msg)


# Global dashboard state
dashboard = DashboardState()


# ============================================================================
# FastAPI Application
# ============================================================================

@asynccontextmanager
async def lifespan(app: FastAPI):
    """Application lifespan manager."""
    logger.info("Dashboard server starting...")
    yield
    logger.info("Dashboard server shutting down...")


app = FastAPI(
    title="Crypto Arbitrage Dashboard",
    description="Real-time visualization of lead-lag arbitrage model",
    lifespan=lifespan,
)

# Mount static files
static_path = Path(__file__).parent / "static"
if static_path.exists():
    app.mount("/static", StaticFiles(directory=static_path), name="static")


@app.get("/", response_class=HTMLResponse)
async def get_dashboard():
    """Serve the main dashboard page."""
    template_path = Path(__file__).parent / "templates" / "dashboard.html"
    if template_path.exists():
        return HTMLResponse(content=template_path.read_text())
    return HTMLResponse(content="<h1>Dashboard template not found</h1>")


@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
    """WebSocket endpoint for real-time updates."""
    await dashboard.manager.connect(websocket)

    try:
        # Send initial state
        await websocket.send_json(dashboard.get_initial_state())

        # Keep connection alive and handle incoming messages
        while True:
            try:
                data = await asyncio.wait_for(
                    websocket.receive_text(),
                    timeout=30.0
                )
                # Handle client messages (e.g., parameter changes)
                msg = json.loads(data)
                if msg.get("type") == "ping":
                    await websocket.send_json({"type": "pong"})
            except asyncio.TimeoutError:
                # Send keepalive
                await websocket.send_json({"type": "keepalive"})
    except WebSocketDisconnect:
        await dashboard.manager.disconnect(websocket)
    except Exception as e:
        logger.error(f"WebSocket error: {e}")
        await dashboard.manager.disconnect(websocket)


@app.get("/api/state")
async def get_state():
    """Get current dashboard state via REST."""
    return dashboard.get_initial_state()


@app.get("/api/price-history/{exchange}/{symbol}")
async def get_price_history(exchange: str, symbol: str):
    """Get price history for a symbol."""
    key = f"{exchange}:{symbol}"
    history = dashboard.price_history.get(key, [])
    return {"exchange": exchange, "symbol": symbol, "history": list(history)}


@app.get("/api/signals")
async def get_signals():
    """Get recent signals."""
    return {"signals": [asdict(s) for s in dashboard.signal_history]}


@app.post("/api/trade-size")
async def update_trade_size(request: dict):
    """
    Update paper trading size.

    Request body:
        {
            "size": 100.0,
            "currency": "USD"  # or "GBP"
        }
    """
    size = request.get("size")
    currency = request.get("currency", "USD")

    if size is None or size <= 0:
        return {"error": "Invalid trade size - must be positive", "status": "error"}

    if currency not in ("USD", "GBP"):
        return {"error": "Currency must be USD or GBP", "status": "error"}

    result = await dashboard.update_trade_size(float(size), currency)
    return result


@app.get("/api/portfolio")
async def get_portfolio():
    """Get current portfolio state."""
    if dashboard._paper_trader_ref:
        state = dashboard._paper_trader_ref.get_portfolio_state()
        state["currency"] = dashboard.trade_currency
        state["trade_size_usd"] = dashboard.trade_size_usd
        total_pnl = state["total_pnl_usd"]
        state["portfolio_value_usd"] = dashboard.trade_size_usd + total_pnl
        return state
    return {"error": "Paper trader not initialized"}


@app.post("/api/fees")
async def update_fees(request: dict):
    """
    Update trading fee for an exchange.

    Request body:
        {
            "exchange": "binance",
            "fee_pct": 0.1  # Fee in percentage (e.g., 0.1 for 0.1%)
        }
    """
    exchange = request.get("exchange")
    fee_pct = request.get("fee_pct")

    if not exchange:
        return {"error": "Missing exchange parameter", "status": "error"}
    if fee_pct is None:
        return {"error": "Missing fee_pct parameter", "status": "error"}

    try:
        fee_pct = float(fee_pct)
    except (ValueError, TypeError):
        return {"error": "Invalid fee_pct - must be a number", "status": "error"}

    result = await dashboard.update_fees(exchange, fee_pct)
    return result


@app.post("/api/spread-threshold")
async def update_spread_threshold(request: dict):
    """
    Update spread threshold in basis points.

    Request body:
        {
            "threshold_bps": 20  # Threshold in basis points
        }
    """
    threshold_bps = request.get("threshold_bps")

    if threshold_bps is None:
        return {"error": "Missing threshold_bps parameter", "status": "error"}

    try:
        threshold_bps = float(threshold_bps)
    except (ValueError, TypeError):
        return {"error": "Invalid threshold_bps - must be a number", "status": "error"}

    result = await dashboard.update_spread_threshold(threshold_bps)
    return result


@app.post("/api/lead-lag-threshold")
async def update_lead_lag_threshold(request: dict):
    """
    Update lead-lag model minimum expected return threshold.

    Request body:
        {
            "min_return_pct": 0.25  # Minimum expected return in percentage
        }
    """
    min_return_pct = request.get("min_return_pct")

    if min_return_pct is None:
        return {"error": "Missing min_return_pct parameter", "status": "error"}

    try:
        min_return_pct = float(min_return_pct)
    except (ValueError, TypeError):
        return {"error": "Invalid min_return_pct - must be a number", "status": "error"}

    result = await dashboard.update_lead_lag_min_return(min_return_pct)
    return result


@app.get("/api/settings")
async def get_settings():
    """Get current fees and threshold settings."""
    return {
        "fees": {k: float(v) * 100 for k, v in dashboard.fees.items()},
        "spread_threshold_bps": float(dashboard.spread_threshold_bps),
        "lead_lag_min_return_pct": dashboard.lead_lag_min_return_pct,
    }


# ============================================================================
# Quantitative Models API
# ============================================================================

@app.get("/api/models")
async def get_all_model_stats():
    """Get statistics from all quantitative models."""
    if dashboard._model_manager is None:
        return {"error": "Model manager not initialized", "models": {}}
    return dashboard._model_manager.get_all_stats()


@app.get("/api/models/{category}")
async def get_category_stats(category: str):
    """Get statistics for a specific model category."""
    if dashboard._model_manager is None:
        return {"error": "Model manager not initialized"}

    valid_categories = ["volatility", "risk", "regime", "jumps", "stat_arb", "microstructure"]
    if category not in valid_categories:
        return {"error": f"Invalid category. Valid: {valid_categories}"}

    return dashboard._model_manager.get_category_stats(category)


@app.get("/api/models/symbol/{symbol}")
async def get_symbol_stats(symbol: str):
    """Get all model statistics for a specific symbol."""
    if dashboard._model_manager is None:
        return {"error": "Model manager not initialized"}
    return dashboard._model_manager.get_symbol_stats(symbol.upper())


@app.get("/quant")
async def get_quant_dashboard():
    """Serve the quantitative models dashboard page."""
    template_path = Path(__file__).parent / "templates" / "quant_models.html"
    if template_path.exists():
        return HTMLResponse(content=template_path.read_text())
    return HTMLResponse(content="<h1>Quantitative models template not found</h1>")


# ============================================================================
# Integration Helpers
# ============================================================================

def get_dashboard() -> DashboardState:
    """Get the global dashboard state for integration."""
    return dashboard


async def run_dashboard_server(host: str = "0.0.0.0", port: int = 8000):
    """Run the dashboard server."""
    import uvicorn
    config = uvicorn.Config(app, host=host, port=port, log_level="info")
    server = uvicorn.Server(config)
    await server.serve()
