"""
Backtesting engine for binary options strategies.

Simulates CALL/PUT trades with fixed payout, tracks:
- Win rate, profit factor, Sharpe ratio
- Max drawdown, consecutive wins/losses
- Equity curve, trade log
"""

import logging
import json
import os
from dataclasses import dataclass, field
from datetime import datetime
from typing import Optional

import numpy as np
import pandas as pd
from config import Config

logger = logging.getLogger(__name__)


@dataclass
class TradeResult:
    timestamp: str
    signal: str
    entry_price: float
    exit_price: float
    result: str  # WIN / LOSS
    confidence: float
    strategy: str
    instrument: str
    reason: str
    pnl: float


@dataclass
class StrategyReport:
    name: str
    instrument: str
    expiry_bars: int = 1
    total_signals: int = 0
    wins: int = 0
    losses: int = 0
    win_rate: float = 0.0
    total_pnl: float = 0.0
    profit_factor: float = 0.0
    max_drawdown_pct: float = 0.0
    max_consecutive_wins: int = 0
    max_consecutive_losses: int = 0
    avg_confidence_win: float = 0.0
    avg_confidence_loss: float = 0.0
    trades: list = field(default_factory=list)
    equity_curve: list = field(default_factory=list)


class StrategyTester:
    def __init__(self):
        self.config = Config()
        self.backtest_config = self.config.BACKTEST
        self.payout = self.backtest_config['payout_pct']
        self.trade_amount = self.backtest_config['trade_amount']

    def test_strategy(self, df: pd.DataFrame, strategy_fn, params: dict,
                      instrument: str, strategy_name: str,
                      expiry_bars: int = 1) -> StrategyReport:
        report = StrategyReport(name=strategy_name, instrument=instrument,
                                expiry_bars=expiry_bars)
        balance = self.backtest_config['initial_balance']
        peak_balance = balance
        max_dd = 0.0
        consecutive_wins = 0
        consecutive_losses = 0
        max_cw = 0
        max_cl = 0
        equity = [balance]

        # Start index: after indicator warmup (max lookback across all strategies)
        max_lookback = 0
        for p in self.config.STRATEGY_PARAMS.values():
            max_lookback = max(max_lookback, p.get('ema_fast', 0), p.get('ema_slow', 0),
                               p.get('bb_period', 0), p.get('div_lookback', 0),
                               p.get('macd_slow', 0), p.get('lookback', 0))
        start_idx = max(max_lookback + 5, 40)

        for i in range(start_idx, len(df)):
            # Check if there's enough data for the expiry horizon
            if i + expiry_bars >= len(df):
                continue

            future_close = df['Close'].iloc[i + expiry_bars]
            current_close = df['Close'].iloc[i]

            if pd.isna(future_close) or pd.isna(current_close) or current_close == 0:
                continue

            result = strategy_fn(df, params, i)
            if result['signal'] is None:
                continue

            report.total_signals += 1

            is_call = result['signal'] == 'CALL'
            is_win = (is_call and future_close > current_close) or (not is_call and future_close < current_close)

            pnl = self.trade_amount * self.payout if is_win else -self.trade_amount
            balance += pnl

            trade = TradeResult(
                timestamp=str(df.index[i]),
                signal=result['signal'],
                entry_price=round(current_close, 5),
                exit_price=round(future_close, 5),
                result='WIN' if is_win else 'LOSS',
                confidence=result['confidence'],
                strategy=strategy_name,
                instrument=instrument,
                reason=result['reason'],
                pnl=round(pnl, 2),
            )
            report.trades.append(trade)

            if is_win:
                report.wins += 1
                consecutive_wins += 1
                consecutive_losses = 0
                max_cw = max(max_cw, consecutive_wins)
            else:
                report.losses += 1
                consecutive_losses += 1
                consecutive_wins = 0
                max_cl = max(max_cl, consecutive_losses)

            peak_balance = max(peak_balance, balance)
            dd = (peak_balance - balance) / peak_balance * 100 if peak_balance > 0 else 0
            max_dd = max(max_dd, dd)
            equity.append(round(balance, 2))

        report.equity_curve = equity
        report.max_consecutive_wins = max_cw
        report.max_consecutive_losses = max_cl

        if report.total_signals > 0:
            report.win_rate = report.wins / report.total_signals * 100

            wins_data = [t for t in report.trades if t.result == 'WIN']
            losses_data = [t for t in report.trades if t.result == 'LOSS']
            report.avg_confidence_win = np.mean([t.confidence for t in wins_data]) if wins_data else 0
            report.avg_confidence_loss = np.mean([t.confidence for t in losses_data]) if losses_data else 0

            gross_profit = sum(t.pnl for t in wins_data)
            gross_loss = abs(sum(t.pnl for t in losses_data))
            report.profit_factor = gross_profit / gross_loss if gross_loss > 0 else float('inf')

        report.total_pnl = round(balance - self.backtest_config['initial_balance'], 2)
        report.max_drawdown_pct = round(max_dd, 2)

        return report

    def run_all(self, data: dict, strategies_map: dict,
                expiry_bars_list: list = None) -> dict:
        """Run all strategies on all instruments for all expiry values.

        Args:
            data: {instrument: DataFrame}
            strategies_map: {name: strategy_fn}
            expiry_bars_list: list of expiry bar counts (e.g. [1, 2, 3, 4, 5]).
                              If None, defaults to [1].

        Returns:
            {f'{instr}_{name}_exp{expiry}': StrategyReport}
        """
        if expiry_bars_list is None:
            expiry_bars_list = [1]

        results = {}
        for instr, df in data.items():
            if df.empty:
                continue
            logger.info("Testing strategies on %s (%d candles)", instr, len(df))
            for expiry in expiry_bars_list:
                for name, fn in strategies_map.items():
                    params = self.config.STRATEGY_PARAMS[name]
                    key = f'{instr}_{name}_exp{expiry}'
                    report = self.test_strategy(df, fn, params, instr, name,
                                                expiry_bars=expiry)
                    results[key] = report
                    self._log_report(report)
        return results

    def _log_report(self, r: StrategyReport):
        status = 'PASS' if r.win_rate >= 55 and r.profit_factor >= 1.0 and r.total_signals >= 30 else 'FAIL'
        logger.info("%s %-18s exp=%-2d | signals=%4d  win=%.1f%%  pf=%.2f  dd=%.1f%%  pnl=%.0f  %s",
                    r.instrument, r.name, r.expiry_bars,
                    r.total_signals, r.win_rate, r.profit_factor,
                    r.max_drawdown_pct, r.total_pnl, status)

    def export_results(self, results: dict, path: str):
        os.makedirs(os.path.dirname(path), exist_ok=True)
        data = {}
        for key, r in results.items():
            data[key] = {
                'expiry_bars': r.expiry_bars,
                'total_signals': r.total_signals,
                'wins': r.wins, 'losses': r.losses,
                'win_rate': r.win_rate,
                'profit_factor': r.profit_factor,
                'max_drawdown_pct': r.max_drawdown_pct,
                'total_pnl': r.total_pnl,
                'max_consecutive_wins': r.max_consecutive_wins,
                'max_consecutive_losses': r.max_consecutive_losses,
                'avg_confidence_win': r.avg_confidence_win,
                'avg_confidence_loss': r.avg_confidence_loss,
                'equity_snapshot': r.equity_curve[-20:] if len(r.equity_curve) > 20 else r.equity_curve,
            }
        with open(path, 'w') as f:
            json.dump(data, f, indent=2, default=str)
        logger.info("Results saved to %s", path)

    def filter_profitable(self, results: dict) -> dict:
        cfg = self.backtest_config
        return {
            k: v for k, v in results.items()
            if v.win_rate >= cfg['min_win_rate'] * 100
            and v.profit_factor >= cfg['min_profit_factor']
            and v.total_signals >= cfg['min_trades']
        }
