"""
Paper trading simulator for tracking hypothetical P&L.

Supports both spread-based arbitrage and lead-lag signals.

Lead-lag positions are tracked as pending until max_hold_ms expires,
then closed at actual market price for realistic P&L calculation.
"""
import asyncio
import time
from dataclasses import dataclass, field
from datetime import datetime
from decimal import Decimal
from typing import Callable, Awaitable

from config import PAPER_TRADE_SIZE_USD, SUMMARY_INTERVAL_S
from models import Tick
from utils import get_logger
from .spread_calculator import SpreadResult
from .lead_lag_brain import LeadLagSignal, SignalDirection


@dataclass
class PaperTrade:
    """Record of a simulated spread arbitrage trade."""
    timestamp: datetime
    symbol: str
    trade_type: str             # "spread" or "lead_lag"
    buy_exchange: str
    sell_exchange: str
    buy_price: Decimal
    sell_price: Decimal
    quantity: Decimal           # Amount of crypto traded
    gross_profit_usd: Decimal   # Before fees
    net_profit_usd: Decimal     # After fees
    spread_bps: Decimal


@dataclass
class LeadLagTrade:
    """Record of a simulated lead-lag trade."""
    timestamp: datetime
    symbol: str
    direction: str              # "long" or "short"
    entry_price: Decimal
    expected_return_pct: float
    actual_return_pct: float    # Set when trade closes
    beta: float
    confidence: float
    pnl_usd: Decimal            # Set when trade closes
    status: str = "open"        # "open", "closed", "expired"
    max_hold_ms: int = 0
    close_timestamp: datetime | None = None
    close_price: Decimal | None = None
    exchange: str = ""          # Exchange where position is held


@dataclass
class PendingPosition:
    """A pending lead-lag position waiting to be closed at actual market price."""
    trade_id: int               # Index in lead_lag_trades
    symbol: str
    exchange: str               # Exchange where we're trading
    direction: str              # "long" or "short"
    entry_price: Decimal
    entry_time_ms: int          # Entry timestamp in ms
    max_hold_ms: int            # Max time to hold
    trade_size_usd: Decimal     # Position size
    expected_return_pct: float
    beta: float
    confidence: float


@dataclass
class TradingStats:
    """Aggregated trading statistics."""
    # Spread arbitrage stats
    spread_trades: int = 0
    spread_gross_profit_usd: Decimal = field(default_factory=lambda: Decimal("0"))
    spread_net_profit_usd: Decimal = field(default_factory=lambda: Decimal("0"))
    spread_volume_usd: Decimal = field(default_factory=lambda: Decimal("0"))
    best_spread_bps: Decimal = field(default_factory=lambda: Decimal("0"))
    worst_spread_bps: Decimal = field(default_factory=lambda: Decimal("9999"))
    avg_spread_bps: Decimal = field(default_factory=lambda: Decimal("0"))

    # Lead-lag stats
    lead_lag_signals: int = 0
    lead_lag_trades_closed: int = 0
    lead_lag_wins: int = 0
    lead_lag_pnl_usd: Decimal = field(default_factory=lambda: Decimal("0"))
    avg_beta: float = 0.0
    avg_confidence: float = 0.0

    # Breakdown
    trades_by_pair: dict = field(default_factory=dict)
    trades_by_exchange: dict = field(default_factory=dict)
    lead_lag_by_symbol: dict = field(default_factory=dict)

    @property
    def total_trades(self) -> int:
        return self.spread_trades + self.lead_lag_signals

    @property
    def total_net_profit_usd(self) -> Decimal:
        return self.spread_net_profit_usd + self.lead_lag_pnl_usd

    @property
    def total_volume_usd(self) -> Decimal:
        return self.spread_volume_usd


class PaperTrader:
    """
    Simulates trade execution and tracks hypothetical P&L.

    Handles both spread-based arbitrage and lead-lag signals.
    Does not place real orders - just logs what would have happened.
    """

    def __init__(
        self,
        trade_size_usd: Decimal | None = None,
        summary_interval_s: float | None = None,
        trade_logger=None,
    ):
        """
        Initialize the paper trader.

        Args:
            trade_size_usd: Hypothetical USD size per trade. Defaults to config value.
            summary_interval_s: Seconds between summary outputs. Defaults to config value.
            trade_logger: Optional TradeLogger instance for detailed trade logging.
        """
        self.trade_size_usd = trade_size_usd or PAPER_TRADE_SIZE_USD
        self.summary_interval_s = summary_interval_s or SUMMARY_INTERVAL_S
        self.trade_logger = trade_logger

        self.spread_trades: list[PaperTrade] = []
        self.lead_lag_trades: list[LeadLagTrade] = []
        self.pending_positions: list[PendingPosition] = []  # Open positions awaiting close
        self.stats = TradingStats()
        self.start_time = datetime.now()
        self._summary_task: asyncio.Task | None = None

        # Track latest prices for position closing
        self._latest_prices: dict[str, dict[str, Decimal]] = {}  # exchange -> symbol -> mid_price

        self.logger = get_logger("paper_trader")

    async def execute(self, signal: SpreadResult) -> None:
        """
        Execute a paper trade based on a spread arbitrage signal.

        Args:
            signal: The arbitrage opportunity to simulate.
        """
        # Calculate trade quantity (in crypto)
        # We buy at ask price, so quantity = USD / ask_price
        quantity = self.trade_size_usd / signal.buy_price

        # Calculate profits
        # Gross: (sell_price - buy_price) * quantity
        gross_profit = (signal.sell_price - signal.buy_price) * quantity

        # Net profit is based on spread_bps (already fee-adjusted)
        # net_profit = trade_size * (spread_bps / 10000)
        net_profit = self.trade_size_usd * (signal.spread_bps / 10000)

        trade = PaperTrade(
            timestamp=datetime.now(),
            symbol=signal.symbol,
            trade_type="spread",
            buy_exchange=signal.buy_exchange,
            sell_exchange=signal.sell_exchange,
            buy_price=signal.buy_price,
            sell_price=signal.sell_price,
            quantity=quantity,
            gross_profit_usd=gross_profit,
            net_profit_usd=net_profit,
            spread_bps=signal.spread_bps,
        )

        self.spread_trades.append(trade)
        self._update_spread_stats(trade)

        self.logger.info(
            f"SPREAD TRADE #{len(self.spread_trades)}: {trade.symbol} "
            f"BUY {trade.quantity:.6f} @ {trade.buy_exchange}:{trade.buy_price:.2f} -> "
            f"SELL @ {trade.sell_exchange}:{trade.sell_price:.2f} "
            f"| Net P&L: ${trade.net_profit_usd:.4f} ({trade.spread_bps:.2f}bps)"
        )

        # Log to trade logger with full mathematical details
        if self.trade_logger:
            from config import FEES, SPREAD_THRESHOLD_BPS
            self.trade_logger.log_spread_trade(
                spread_result=signal,
                trade_size_usd=self.trade_size_usd,
                quantity=quantity,
                gross_profit_usd=gross_profit,
                net_profit_usd=net_profit,
                portfolio_state=self.get_portfolio_state(),
                fees=FEES,
                threshold_bps=SPREAD_THRESHOLD_BPS,
            )

    async def execute_lead_lag(
        self,
        signal: LeadLagSignal,
        entry_price: Decimal | None = None,
        exchange: str = "",
    ) -> None:
        """
        Open a lead-lag position for paper trading.

        Unlike the old implementation that estimated P&L at entry, this now
        opens a real position that will be closed at actual market price
        when max_hold_ms expires (via on_tick).

        Args:
            signal: The lead-lag opportunity signal.
            entry_price: Current market price for entry (mid price).
            exchange: Exchange where position is opened.
        """
        direction = "long" if signal.direction == SignalDirection.LONG else "short"

        # Use provided entry price or try to get from latest prices
        if entry_price is None:
            exchange_prices = self._latest_prices.get(exchange, {})
            entry_price = exchange_prices.get(signal.lagger_symbol, Decimal("0"))

        if entry_price == Decimal("0"):
            self.logger.warning(
                f"No entry price available for {signal.lagger_symbol} on {exchange}, "
                f"cannot open position"
            )
            return

        trade_id = len(self.lead_lag_trades)

        # Record the trade as "open"
        trade = LeadLagTrade(
            timestamp=datetime.now(),
            symbol=signal.lagger_symbol,
            direction=direction,
            entry_price=entry_price,
            expected_return_pct=signal.expected_return_pct,
            actual_return_pct=0.0,  # Set when closed
            beta=signal.rolling_beta,
            confidence=signal.confidence,
            pnl_usd=Decimal("0"),  # Set when closed
            status="open",
            max_hold_ms=signal.max_hold_ms,
            close_timestamp=None,
            close_price=None,
            exchange=exchange,
        )
        self.lead_lag_trades.append(trade)

        # Create pending position for tracking
        position = PendingPosition(
            trade_id=trade_id,
            symbol=signal.lagger_symbol,
            exchange=exchange,
            direction=direction,
            entry_price=entry_price,
            entry_time_ms=int(time.time() * 1000),
            max_hold_ms=signal.max_hold_ms,
            trade_size_usd=self.trade_size_usd,
            expected_return_pct=signal.expected_return_pct,
            beta=signal.rolling_beta,
            confidence=signal.confidence,
        )
        self.pending_positions.append(position)

        # Update signal count (P&L updated when position closes)
        self.stats.lead_lag_signals += 1

        self.logger.info(
            f"LEAD-LAG OPEN #{trade_id + 1}: {direction.upper()} {trade.symbol} "
            f"@ {exchange}:{entry_price:.2f} "
            f"| β={trade.beta:.2f} conf={trade.confidence:.2f} "
            f"| Expected={signal.expected_return_pct:.3f}% "
            f"| MaxHold={signal.max_hold_ms}ms"
        )

        # Log to trade logger with full mathematical details
        if self.trade_logger:
            from config import (
                LEAD_LAG_LEADER_Z_THRESHOLD,
                LEAD_LAG_LAG_Z_THRESHOLD,
                LEAD_LAG_GAP_THRESHOLD,
                LEAD_LAG_MIN_CORRELATION,
                LEAD_LAG_MIN_CONFIDENCE,
                LEAD_LAG_MAX_LAG_MS,
            )
            config_dict = {
                "leader_z_threshold": LEAD_LAG_LEADER_Z_THRESHOLD,
                "lag_z_threshold": LEAD_LAG_LAG_Z_THRESHOLD,
                "gap_threshold": LEAD_LAG_GAP_THRESHOLD,
                "min_correlation": LEAD_LAG_MIN_CORRELATION,
                "min_confidence": LEAD_LAG_MIN_CONFIDENCE,
                "max_lag_ms": LEAD_LAG_MAX_LAG_MS,
            }
            # Build stats dict from signal data (Bug Fix #1: now uses actual signal fields)
            stats_dict = {
                "idiosyncratic_std": signal.idiosyncratic_std,
                "total_std": signal.total_std,
            }
            self.trade_logger.log_lead_lag_open(
                signal=signal,
                trade_size_usd=self.trade_size_usd,
                portfolio_state=self.get_portfolio_state(),
                stats=stats_dict,
                config=config_dict,
            )

    async def on_tick(self, tick: Tick) -> None:
        """
        Process incoming ticks to:
        1. Update latest prices for position closing
        2. Check if any pending positions have expired and close them

        This is the key fix - positions are closed at ACTUAL market price,
        not at estimated P&L from entry time.

        Args:
            tick: Incoming market tick.
        """
        # Update latest prices
        if tick.exchange not in self._latest_prices:
            self._latest_prices[tick.exchange] = {}
        self._latest_prices[tick.exchange][tick.symbol] = tick.mid_price()

        # Check pending positions for expiry
        current_time_ms = int(time.time() * 1000)
        positions_to_close: list[PendingPosition] = []

        for position in self.pending_positions:
            # Check if position has exceeded max hold time
            elapsed_ms = current_time_ms - position.entry_time_ms
            if elapsed_ms >= position.max_hold_ms:
                positions_to_close.append(position)

        # Close expired positions
        for position in positions_to_close:
            await self._close_position(position, current_time_ms)

    async def _close_position(self, position: PendingPosition, close_time_ms: int) -> None:
        """
        Close a pending position at actual market price.

        Calculates real P&L based on price movement, not estimated return.
        """
        # Get current price for this symbol on the exchange
        exchange_prices = self._latest_prices.get(position.exchange, {})
        close_price = exchange_prices.get(position.symbol)

        if close_price is None:
            # Fallback: check any exchange for this symbol
            for exch, prices in self._latest_prices.items():
                if position.symbol in prices:
                    close_price = prices[position.symbol]
                    break

        if close_price is None:
            self.logger.warning(
                f"No close price available for {position.symbol}, marking expired"
            )
            # Update trade status
            if position.trade_id < len(self.lead_lag_trades):
                trade = self.lead_lag_trades[position.trade_id]
                trade.status = "expired"
                trade.close_timestamp = datetime.now()
            self.pending_positions.remove(position)
            return

        # Calculate actual P&L
        entry = position.entry_price
        exit_price = close_price

        if position.direction == "long":
            # Long: profit if price went up
            return_pct = float((exit_price - entry) / entry * 100)
            pnl = position.trade_size_usd * (exit_price - entry) / entry
        else:
            # Short: profit if price went down
            return_pct = float((entry - exit_price) / entry * 100)
            pnl = position.trade_size_usd * (entry - exit_price) / entry

        # Update the trade record
        if position.trade_id < len(self.lead_lag_trades):
            trade = self.lead_lag_trades[position.trade_id]
            trade.status = "closed"
            trade.close_timestamp = datetime.now()
            trade.close_price = close_price
            trade.actual_return_pct = return_pct
            trade.pnl_usd = pnl

            # Update stats with actual P&L
            self._finalize_lead_lag_stats(trade, position)

            self.logger.info(
                f"LEAD-LAG CLOSE #{position.trade_id + 1}: {position.direction.upper()} {position.symbol} "
                f"| Entry={entry:.2f} Exit={exit_price:.2f} "
                f"| Return={return_pct:.3f}% (expected={position.expected_return_pct:.3f}%) "
                f"| P&L: ${pnl:.4f}"
            )

            # Log to trade logger with complete P&L analysis
            if self.trade_logger:
                position_duration_ms = close_time_ms - position.entry_time_ms
                self.trade_logger.log_lead_lag_close(
                    trade_id=position.trade_id + 1,
                    symbol=position.symbol,
                    exchange=position.exchange,
                    direction=position.direction.upper(),
                    entry_price=float(position.entry_price),
                    close_price=float(close_price),
                    trade_size_usd=float(position.trade_size_usd),
                    expected_return_pct=position.expected_return_pct,
                    actual_return_pct=return_pct,
                    pnl_usd=float(pnl),
                    position_duration_ms=position_duration_ms,
                    portfolio_state=self.get_portfolio_state(),
                )

        # Remove from pending
        self.pending_positions.remove(position)

    def _finalize_lead_lag_stats(self, trade: LeadLagTrade, position: PendingPosition) -> None:
        """Update stats when a lead-lag position is closed with actual P&L."""
        self.stats.lead_lag_trades_closed += 1
        self.stats.lead_lag_pnl_usd += trade.pnl_usd

        if trade.pnl_usd > 0:
            self.stats.lead_lag_wins += 1

        # Update running averages
        n = self.stats.lead_lag_trades_closed
        self.stats.avg_beta = (self.stats.avg_beta * (n - 1) + trade.beta) / n
        self.stats.avg_confidence = (self.stats.avg_confidence * (n - 1) + trade.confidence) / n

        # Track by symbol
        if trade.symbol not in self.stats.lead_lag_by_symbol:
            self.stats.lead_lag_by_symbol[trade.symbol] = {
                "count": 0,
                "profit": Decimal("0"),
                "avg_beta": 0.0,
            }
        sym_stats = self.stats.lead_lag_by_symbol[trade.symbol]
        sym_stats["count"] += 1
        sym_stats["profit"] += trade.pnl_usd
        sym_stats["avg_beta"] = (
            (sym_stats["avg_beta"] * (sym_stats["count"] - 1) + trade.beta)
            / sym_stats["count"]
        )

    def _update_spread_stats(self, trade: PaperTrade) -> None:
        """Update aggregated statistics with new spread trade."""
        self.stats.spread_trades += 1
        self.stats.spread_gross_profit_usd += trade.gross_profit_usd
        self.stats.spread_net_profit_usd += trade.net_profit_usd
        self.stats.spread_volume_usd += self.trade_size_usd

        if trade.spread_bps > self.stats.best_spread_bps:
            self.stats.best_spread_bps = trade.spread_bps
        if trade.spread_bps < self.stats.worst_spread_bps:
            self.stats.worst_spread_bps = trade.spread_bps

        # Running average
        n = self.stats.spread_trades
        self.stats.avg_spread_bps = (
            (self.stats.avg_spread_bps * (n - 1) + trade.spread_bps) / n
        )

        # Track by pair
        pair_key = trade.symbol
        if pair_key not in self.stats.trades_by_pair:
            self.stats.trades_by_pair[pair_key] = {"count": 0, "profit": Decimal("0")}
        self.stats.trades_by_pair[pair_key]["count"] += 1
        self.stats.trades_by_pair[pair_key]["profit"] += trade.net_profit_usd

        # Track by exchange pair
        exchange_key = f"{trade.buy_exchange}->{trade.sell_exchange}"
        if exchange_key not in self.stats.trades_by_exchange:
            self.stats.trades_by_exchange[exchange_key] = {"count": 0, "profit": Decimal("0")}
        self.stats.trades_by_exchange[exchange_key]["count"] += 1
        self.stats.trades_by_exchange[exchange_key]["profit"] += trade.net_profit_usd


    def get_summary(self) -> str:
        """Generate a summary string of trading performance."""
        runtime = datetime.now() - self.start_time
        hours = runtime.total_seconds() / 3600

        lines = [
            "=" * 70,
            f"PAPER TRADING SUMMARY (Runtime: {runtime})",
            "=" * 70,
        ]

        # Overall stats
        total_pnl = self.stats.spread_net_profit_usd + self.stats.lead_lag_pnl_usd
        lines.append(f"Total P&L:         ${total_pnl:,.4f}")
        lines.append("")

        # Spread arbitrage section
        lines.append("-" * 35 + " SPREAD ARBITRAGE " + "-" * 17)
        lines.append(f"  Trades:          {self.stats.spread_trades}")
        lines.append(f"  Volume:          ${self.stats.spread_volume_usd:,.2f}")
        lines.append(f"  Net P&L:         ${self.stats.spread_net_profit_usd:,.4f}")
        if self.stats.spread_trades > 0:
            lines.append(f"  Avg Spread:      {self.stats.avg_spread_bps:.2f} bps")
            lines.append(f"  Best Trade:      {self.stats.best_spread_bps:.2f} bps")
            lines.append(f"  Worst Trade:     {self.stats.worst_spread_bps:.2f} bps")

        # Lead-lag section
        lines.append("")
        lines.append("-" * 35 + " LEAD-LAG MODEL " + "-" * 19)
        lines.append(f"  Signals:         {self.stats.lead_lag_signals}")
        lines.append(f"  Positions Open:  {len(self.pending_positions)}")
        lines.append(f"  Closed:          {self.stats.lead_lag_trades_closed}")
        lines.append(f"  Actual P&L:      ${self.stats.lead_lag_pnl_usd:,.4f}")
        if self.stats.lead_lag_trades_closed > 0:
            win_rate = self.stats.lead_lag_wins / self.stats.lead_lag_trades_closed * 100
            lines.append(f"  Win Rate:        {win_rate:.1f}%")
            lines.append(f"  Avg Beta:        {self.stats.avg_beta:.2f}")
            lines.append(f"  Avg Confidence:  {self.stats.avg_confidence:.2f}")

        # Hourly rates
        if hours > 0:
            lines.append("")
            lines.append("-" * 35 + " HOURLY RATES " + "-" * 21)
            total_trades = self.stats.spread_trades + self.stats.lead_lag_signals
            trades_per_hour = total_trades / hours
            profit_per_hour = float(total_pnl) / hours
            lines.append(f"  Trades/Hour:     {trades_per_hour:.1f}")
            lines.append(f"  Profit/Hour:     ${profit_per_hour:.4f}")

        # Breakdown by symbol
        if self.stats.trades_by_pair or self.stats.lead_lag_by_symbol:
            lines.append("")
            lines.append("-" * 35 + " BY SYMBOL " + "-" * 24)

            # Combine spread and lead-lag by symbol
            all_symbols = set(self.stats.trades_by_pair.keys()) | set(self.stats.lead_lag_by_symbol.keys())
            for symbol in sorted(all_symbols):
                spread_data = self.stats.trades_by_pair.get(symbol, {"count": 0, "profit": Decimal("0")})
                ll_data = self.stats.lead_lag_by_symbol.get(symbol, {"count": 0, "profit": Decimal("0"), "avg_beta": 0})

                total_count = spread_data["count"] + ll_data["count"]
                total_profit = spread_data["profit"] + ll_data["profit"]

                parts = [f"  {symbol}: {total_count} trades, ${total_profit:.4f}"]
                if ll_data["count"] > 0:
                    parts.append(f"(β={ll_data['avg_beta']:.2f})")
                lines.append(" ".join(parts))

        # Breakdown by exchange route (spread only)
        if self.stats.trades_by_exchange:
            lines.append("")
            lines.append("-" * 35 + " BY EXCHANGE " + "-" * 22)
            for route, data in sorted(self.stats.trades_by_exchange.items()):
                lines.append(f"  {route}: {data['count']} trades, ${data['profit']:.4f}")

        lines.append("=" * 70)
        return "\n".join(lines)

    def print_summary(self) -> None:
        """Print the trading summary to the log."""
        self.logger.info("\n" + self.get_summary())

    async def start_summary_loop(self) -> None:
        """Start a background task that prints summaries periodically."""
        self._summary_task = asyncio.create_task(self._summary_loop())

    async def _summary_loop(self) -> None:
        """Periodically print trading summaries."""
        while True:
            await asyncio.sleep(self.summary_interval_s)
            total_trades = self.stats.spread_trades + self.stats.lead_lag_signals
            if total_trades > 0:
                self.print_summary()

    async def stop(self) -> None:
        """Stop the summary loop and print final summary."""
        if self._summary_task:
            self._summary_task.cancel()
            try:
                await self._summary_task
            except asyncio.CancelledError:
                pass
        self.print_summary()

    def set_trade_size(self, new_size_usd: Decimal, currency: str = "USD") -> dict:
        """
        Update the paper trading position size.

        Args:
            new_size_usd: New trade size in USD (or converted to USD if GBP)
            currency: Original currency for logging purposes

        Returns:
            Dict with confirmation and current portfolio state
        """
        old_size = self.trade_size_usd
        self.trade_size_usd = new_size_usd

        self.logger.info(
            f"Trade size updated: ${old_size} -> ${new_size_usd} ({currency})"
        )

        return {
            "old_size_usd": float(old_size),
            "new_size_usd": float(new_size_usd),
            "currency": currency,
            "pending_positions": len(self.pending_positions),
        }

    def get_portfolio_state(self) -> dict:
        """
        Return current portfolio state for dashboard display.

        Returns:
            Dict with all portfolio metrics for real-time display
        """
        return {
            "trade_size_usd": float(self.trade_size_usd),
            "spread_trades": self.stats.spread_trades,
            "spread_pnl_usd": float(self.stats.spread_net_profit_usd),
            "lead_lag_signals": self.stats.lead_lag_signals,
            "lead_lag_trades_closed": self.stats.lead_lag_trades_closed,
            "lead_lag_wins": self.stats.lead_lag_wins,
            "lead_lag_pnl_usd": float(self.stats.lead_lag_pnl_usd),
            "total_pnl_usd": float(self.stats.total_net_profit_usd),
            "pending_positions": len(self.pending_positions),
            "total_volume_usd": float(self.stats.total_volume_usd),
            "win_rate": (
                self.stats.lead_lag_wins / self.stats.lead_lag_trades_closed * 100
                if self.stats.lead_lag_trades_closed > 0 else 0.0
            ),
        }

    # Legacy properties for backwards compatibility
    @property
    def trades(self) -> list[PaperTrade]:
        """All spread trades (for backwards compatibility)."""
        return self.spread_trades
