"""Walk-Forward Backtesting Engine + Monte Carlo Validation.

Evaluates ML model and signal performance on historical data using
expanding-window methodology (no lookahead bias).
"""
from __future__ import annotations
import logging
import numpy as np
import pandas as pd
from dataclasses import dataclass, field
from typing import Optional

logger = logging.getLogger(__name__)


@dataclass
class Trade:
    """Single trade record."""
    entry_date: str
    exit_date: str
    direction: str  # 'long' or 'short'
    entry_price: float
    exit_price: float
    return_pct: float
    pnl: float
    was_correct: bool
    signal_score: float = 0.0
    confidence: str = ''


@dataclass
class BacktestResults:
    """Comprehensive backtest results."""
    trades: list[Trade] = field(default_factory=list)
    total_trades: int = 0
    winning_trades: int = 0
    losing_trades: int = 0
    winrate: float = 0.0
    avg_return: float = 0.0
    avg_win: float = 0.0
    avg_loss: float = 0.0
    profit_factor: float = 0.0
    sharpe_ratio: float = 0.0
    max_drawdown: float = 0.0
    total_return: float = 0.0
    equity_curve: list[float] = field(default_factory=list)
    monthly_returns: dict[str, float] = field(default_factory=dict)

    # Monte Carlo results
    mc_winrate_ci: tuple[float, float] = (0.0, 0.0)
    mc_sharpe_ci: tuple[float, float] = (0.0, 0.0)
    mc_worst_drawdown: float = 0.0

    def summary(self) -> dict:
        return {
            'total_trades': self.total_trades,
            'winrate': round(self.winrate * 100, 1),
            'avg_return': round(self.avg_return * 100, 2),
            'profit_factor': round(self.profit_factor, 2),
            'sharpe_ratio': round(self.sharpe_ratio, 2),
            'max_drawdown': round(self.max_drawdown * 100, 1),
            'total_return': round(self.total_return * 100, 1),
            'mc_winrate_95ci': [round(x * 100, 1) for x in self.mc_winrate_ci],
            'mc_sharpe_95ci': [round(x, 2) for x in self.mc_sharpe_ci],
        }


class WalkForwardBacktester:
    """Walk-forward backtesting with expanding training window.

    Methodology:
    1. Start with initial training window
    2. Train model on [0:t]
    3. Predict at t+1
    4. Record trade result
    5. Expand window: train on [0:t+1], predict at t+2
    6. Repeat until end of data
    """

    def __init__(
        self,
        initial_train_ratio: float = 0.6,
        fee_pct: float = 0.62,  # Roundtrip fee (buy + sell)
        min_score_threshold: float = 30.0,
    ):
        self.initial_train_ratio = initial_train_ratio
        self.fee_pct = fee_pct
        self.min_score_threshold = min_score_threshold

    def backtest(
        self, df: pd.DataFrame, scoring_fn=None,
        timeframe: str = '1D',
    ) -> BacktestResults:
        """Run walk-forward backtest.

        Args:
            df: Full OHLCV DataFrame (chronologically sorted)
            scoring_fn: Function that takes (feature_df, closes) → dict with
                       'signal', 'score', 'confidence'
            timeframe: Timeframe for context

        Returns:
            BacktestResults with full metrics
        """
        if df is None or len(df) < 100:
            return BacktestResults()

        if scoring_fn is None:
            scoring_fn = self._default_scoring_fn

        n = len(df)
        train_end = int(n * self.initial_train_ratio)
        trades = []
        equity = [1.0]

        for t in range(train_end, n - 1):
            train_df = df.iloc[:t]
            test_row = df.iloc[t]
            next_row = df.iloc[t + 1]

            try:
                result = scoring_fn(train_df, train_df['close'].values)
            except Exception:
                continue

            signal = result.get('signal', 'HOLD')
            score = result.get('score', 0)

            if signal == 'HOLD' or abs(score) < self.min_score_threshold:
                equity.append(equity[-1])
                continue

            entry_price = float(test_row['close'])
            exit_price = float(next_row['close'])

            if signal == 'BUY':
                raw_return = (exit_price - entry_price) / entry_price
            else:  # SELL
                raw_return = (entry_price - exit_price) / entry_price

            net_return = raw_return - (self.fee_pct / 100)
            was_correct = raw_return > 0

            trade = Trade(
                entry_date=str(test_row.name) if hasattr(test_row, 'name') else str(t),
                exit_date=str(next_row.name) if hasattr(next_row, 'name') else str(t + 1),
                direction='long' if signal == 'BUY' else 'short',
                entry_price=entry_price,
                exit_price=exit_price,
                return_pct=net_return,
                pnl=net_return * equity[-1],
                was_correct=was_correct,
                signal_score=score,
                confidence=result.get('confidence', ''),
            )
            trades.append(trade)
            equity.append(equity[-1] * (1 + net_return))

        return self._compute_results(trades, equity)

    def _default_scoring_fn(self, train_df, closes):
        """Simple scoring based on recent momentum."""
        if len(closes) < 20:
            return {'signal': 'HOLD', 'score': 0}

        sma_short = np.mean(closes[-5:])
        sma_long = np.mean(closes[-20:])
        score = (sma_short / sma_long - 1) * 100 * 20  # Scale to ~30 range

        if score > 30:
            return {'signal': 'BUY', 'score': score}
        elif score < -30:
            return {'signal': 'SELL', 'score': score}
        return {'signal': 'HOLD', 'score': score}

    def _compute_results(
        self, trades: list[Trade], equity: list[float]
    ) -> BacktestResults:
        """Compute comprehensive metrics from trades."""
        results = BacktestResults(trades=trades, equity_curve=equity)

        if not trades:
            return results

        returns = [t.return_pct for t in trades]
        wins = [r for r in returns if r > 0]
        losses = [r for r in returns if r <= 0]

        results.total_trades = len(trades)
        results.winning_trades = len(wins)
        results.losing_trades = len(losses)
        results.winrate = len(wins) / len(trades)
        results.avg_return = float(np.mean(returns))
        results.avg_win = float(np.mean(wins)) if wins else 0.0
        results.avg_loss = float(np.mean(losses)) if losses else 0.0

        # Profit factor
        gross_profit = sum(wins)
        gross_loss = abs(sum(losses))
        results.profit_factor = gross_profit / gross_loss if gross_loss > 0 else float('inf')

        # Sharpe ratio (annualized, assuming daily)
        if len(returns) > 1 and np.std(returns) > 0:
            results.sharpe_ratio = float(
                np.mean(returns) / np.std(returns) * np.sqrt(252)
            )

        # Max drawdown
        peak = equity[0]
        max_dd = 0.0
        for val in equity:
            if val > peak:
                peak = val
            dd = (peak - val) / peak
            if dd > max_dd:
                max_dd = dd
        results.max_drawdown = max_dd

        # Total return
        results.total_return = (equity[-1] / equity[0]) - 1

        # Monte Carlo validation
        mc = self._monte_carlo(returns)
        results.mc_winrate_ci = mc['winrate_ci']
        results.mc_sharpe_ci = mc['sharpe_ci']
        results.mc_worst_drawdown = mc['worst_drawdown']

        return results

    def _monte_carlo(
        self, returns: list[float], n_simulations: int = 1000
    ) -> dict:
        """Monte Carlo simulation for confidence intervals.

        Randomly reorders trades to test if results are robust
        or dependent on specific trade sequence.
        """
        if len(returns) < 10:
            return {
                'winrate_ci': (0.0, 0.0),
                'sharpe_ci': (0.0, 0.0),
                'worst_drawdown': 0.0,
            }

        returns_arr = np.array(returns)
        winrates = []
        sharpes = []
        max_drawdowns = []

        for _ in range(n_simulations):
            # Random permutation of trade returns
            perm = np.random.permutation(returns_arr)

            # Winrate
            winrates.append(np.mean(perm > 0))

            # Sharpe
            if np.std(perm) > 0:
                sharpes.append(np.mean(perm) / np.std(perm) * np.sqrt(252))
            else:
                sharpes.append(0)

            # Max drawdown
            equity = np.cumprod(1 + perm)
            peak = np.maximum.accumulate(equity)
            dd = (peak - equity) / peak
            max_drawdowns.append(float(np.max(dd)))

        return {
            'winrate_ci': (
                float(np.percentile(winrates, 2.5)),
                float(np.percentile(winrates, 97.5)),
            ),
            'sharpe_ci': (
                float(np.percentile(sharpes, 2.5)),
                float(np.percentile(sharpes, 97.5)),
            ),
            'worst_drawdown': float(np.percentile(max_drawdowns, 95)),
        }
