#!/usr/bin/env python3
"""Hybrid backtest: MACD+RSI → NN quality filtering threshold sensitivity.

Tests different nn_entry_threshold values (0.3-0.7) on OOS data.
For each ticker:
  1. MACD+RSI strategy finds candidate bars
  2. Per-ticker NN model evaluates signal quality
  3. Trades filtered by quality threshold, using NN-predicted SL/TP
  4. Reports P&L, WR, PF at each threshold
"""

from __future__ import annotations

import sys
import json
import time
from datetime import datetime, date
from pathlib import Path
from typing import Dict, List, Optional, Tuple

import numpy as np
import pandas as pd
from loguru import logger

sys.path.insert(0, str(Path(__file__).parent.parent))

from core.data_loader import DataPreparator
from ai.strategy_signals import get_strategy_signals, StrategyConfig
from ai.inference_v2 import NeuralPredictor
from ai.features import calculate_atr


MOEX_TICKERS = [
    "SBER", "NVTK", "PHOR", "ROSN", "VTBR", "LKOH", "ASTR",
]

THRESHOLDS = [0.3, 0.4, 0.5, 0.6, 0.7]

MODELS_DIR = Path(__file__).parent.parent / "models"
RESULTS_DIR = Path(__file__).parent.parent / "reports"


def load_per_ticker_predictor(ticker: str) -> Tuple[Optional[NeuralPredictor], str]:
    """Load per-ticker model or fallback to strategy model."""
    # Try per-ticker model first
    per_ticker_path = MODELS_DIR / f"{ticker}_strategy.pt"
    if per_ticker_path.exists():
        logger.info(f"[{ticker}] Using per-ticker model: {per_ticker_path.name}")
        predictor = NeuralPredictor(
            model_path=per_ticker_path,
            entry_threshold=0.0,  # No threshold — we filter by nn_entry_threshold in scanner
            min_confidence=0.0,
        )
        return predictor, per_ticker_path.name

    # Fallback to strategy model
    fallback_path = MODELS_DIR / "neural_trader_strategy.pt"
    if fallback_path.exists():
        logger.info(f"[{ticker}] No per-ticker model, fallback to {fallback_path.name}")
        predictor = NeuralPredictor(
            model_path=fallback_path,
            entry_threshold=0.0,
            min_confidence=0.0,
        )
        return predictor, fallback_path.name

    return None, "no_model"


def run_hybrid_backtest(
    ticker: str,
    start: str,
    end: str,
    thresholds: List[float],
) -> Dict:
    """Run hybrid backtest for one ticker at all thresholds.

    Returns dict with ticker info and per-threshold results.
    """
    start_ts = int(datetime.strptime(start, "%Y-%m-%d").timestamp())
    end_ts = int(datetime.strptime(end, "%Y-%m-%d").timestamp())

    try:
        # Load data
        prep = DataPreparator([ticker], ["H1"])
        df = prep.load_h1(ticker, start_ts, end_ts)

        if df.empty or len(df) < 50:
            return {"ticker": ticker, "error": "No data", "n_bars": 0}

        df = df.sort_values("timestamp").reset_index(drop=True)
        n_bars = len(df)
        logger.info(f"[{ticker}] Loaded {n_bars} bars ({start} → {end})")

        # Stage 1: MACD+RSI strategy signals
        config = StrategyConfig()
        signals_df = get_strategy_signals(df, config)

        if signals_df.empty:
            return {"ticker": ticker, "error": "No strategy signals", "n_bars": n_bars}

        # Merge signals into df
        df["strategy_signal"] = signals_df["strategy_signal"].values
        df["strategy_side"] = signals_df["side"].values

        # Find strategy signal bars
        signal_indices = df[df["strategy_signal"] == 1].index.tolist()
        n_signals = len(signal_indices)
        logger.info(f"[{ticker}] {n_signals} MACD+RSI signal bars ({n_signals/n_bars*100:.1f}%)")

        # Load NN predictor
        predictor, model_name = load_per_ticker_predictor(ticker)
        if predictor is None:
            return {"ticker": ticker, "error": "No NN model", "n_bars": n_bars}

        # Stage 2: NN evaluation for each signal bar
        nn_results = []
        atr = calculate_atr(df, 14)

        for i in signal_indices:
            side = df.iloc[i]["strategy_side"]
            if not side or side not in ("LONG", "SHORT"):
                continue

            nn_result = predictor.evaluate_strategy_signal(df, i, side)

            if nn_result is None:
                continue

            entry_price = df.iloc[i]["Close"]
            atr_val = atr.iloc[i] if i < len(atr) else atr.iloc[-1]
            sl_atr = nn_result["sl_distance_atr"]
            tp_atr = nn_result["tp_distance_atr"]

            if side == "LONG":
                sl_price = entry_price - atr_val * sl_atr
                tp_price = entry_price + atr_val * tp_atr
            else:
                sl_price = entry_price + atr_val * sl_atr
                tp_price = entry_price - atr_val * tp_atr

            nn_results.append({
                "bar_idx": i,
                "timestamp": df.iloc[i]["timestamp"],
                "side": side,
                "entry_price": entry_price,
                "sl_price": sl_price,
                "tp_price": tp_price,
                "nn_proba": nn_result["entry_proba"],
                "nn_confidence": nn_result["confidence"],
                "sl_atr": sl_atr,
                "tp_atr": tp_atr,
            })

        if not nn_results:
            return {"ticker": ticker, "error": "NN evaluated no signals", "n_bars": n_bars}

        logger.info(f"[{ticker}] NN evaluated {len(nn_results)}/{n_signals} signals")

        # Backtest for each threshold
        threshold_results = {}

        for thr in thresholds:
            # Filter by quality threshold
            filtered = [r for r in nn_results if r["nn_proba"] >= thr]
            n_filtered = len(filtered)

            if n_filtered == 0:
                threshold_results[str(thr)] = {
                    "threshold": thr,
                    "signals": len(nn_results),
                    "accepted": 0,
                    "trades": 0,
                    "win_rate": 0,
                    "profit_factor": 0,
                    "total_pnl": 0,
                    "sharpe": 0,
                    "max_drawdown": 0,
                }
                continue

            # Run backtest with NN-predicted SL/TP
            trades = []
            equity = [1_000_000]

            # Sort by bar index to process chronologically
            filtered_sorted = sorted(filtered, key=lambda x: x["bar_idx"])
            in_position = False
            entry_price = 0.0
            entry_bar = -1
            side = ""
            sl_price = 0.0
            tp_price = 0.0

            for i in range(len(df)):
                if in_position:
                    high = df.iloc[i]["High"]
                    low = df.iloc[i]["Low"]
                    close = df.iloc[i]["Close"]

                    if side == "LONG":
                        if high >= tp_price:
                            pnl = (tp_price - entry_price) / entry_price
                            trades.append(pnl)
                            equity.append(equity[-1] * (1 + pnl))
                            in_position = False
                        elif low <= sl_price:
                            pnl = (sl_price - entry_price) / entry_price
                            trades.append(pnl)
                            equity.append(equity[-1] * (1 + pnl))
                            in_position = False
                    elif side == "SHORT":
                        if low <= tp_price:
                            pnl = (entry_price - tp_price) / entry_price
                            trades.append(pnl)
                            equity.append(equity[-1] * (1 + pnl))
                            in_position = False
                        elif high >= sl_price:
                            pnl = (entry_price - sl_price) / entry_price
                            trades.append(pnl)
                            equity.append(equity[-1] * (1 + pnl))
                            in_position = False

                    # Force close after 6 bars (same as strategy)
                    if in_position and (i - entry_bar) >= config.max_holding_bars:
                        close_pnl = (close - entry_price) / entry_price if side == "LONG" else (entry_price - close) / entry_price
                        trades.append(close_pnl)
                        equity.append(equity[-1] * (1 + close_pnl))
                        in_position = False

                # Check for new signal at this bar
                signals_at_bar = [r for r in filtered_sorted if r["bar_idx"] == i]
                if not in_position and signals_at_bar:
                    sig = signals_at_bar[0]
                    entry_price = sig["entry_price"]
                    side = sig["side"]
                    entry_bar = i
                    sl_price = sig["sl_price"]
                    tp_price = sig["tp_price"]
                    in_position = True

            if not trades:
                threshold_results[str(thr)] = {
                    "threshold": thr,
                    "signals": len(nn_results),
                    "accepted": n_filtered,
                    "trades": 0,
                    "win_rate": 0,
                    "profit_factor": 0,
                    "total_pnl": 0,
                    "sharpe": 0,
                    "max_drawdown": 0,
                }
                continue

            trades_arr = np.array(trades)
            wins = trades_arr > 0
            losses = trades_arr <= 0
            win_rate = float(wins.mean() * 100)
            total_pnl = float((equity[-1] / 1_000_000 - 1) * 1_000_000)
            avg_win = float(trades_arr[wins].mean()) if wins.any() else 0
            avg_loss = float(abs(trades_arr[losses].mean())) if losses.any() else 0
            profit_factor = float(trades_arr[wins].sum() / abs(trades_arr[losses].sum())) if losses.any() and abs(trades_arr[losses].sum()) > 1e-10 else 0
            sharpe = float(trades_arr.mean() / trades_arr.std() * np.sqrt(252)) if trades_arr.std() > 0 else 0

            # Max drawdown
            equity_arr = np.array(equity)
            peak = np.maximum.accumulate(equity_arr)
            drawdown = (peak - equity_arr) / peak * 100
            max_dd = float(drawdown.max())

            threshold_results[str(thr)] = {
                "threshold": thr,
                "signals": len(nn_results),
                "accepted": n_filtered,
                "trades": len(trades),
                "win_rate": round(win_rate, 1),
                "profit_factor": round(profit_factor, 2),
                "total_pnl": round(total_pnl, 0),
                "sharpe": round(sharpe, 2),
                "max_drawdown": round(max_dd, 0),
                "avg_win_pct": round(avg_win * 100, 2),
                "avg_loss_pct": round(avg_loss * 100, 2),
            }

        return {
            "ticker": ticker,
            "n_bars": n_bars,
            "n_strategy_signals": n_signals,
            "n_nn_evaluated": len(nn_results),
            "model": model_name,
            "thresholds": threshold_results,
        }

    except Exception as e:
        logger.error(f"[{ticker}] Error: {e}")
        import traceback
        traceback.print_exc()
        return {"ticker": ticker, "error": str(e)}


def print_summary(all_results: List[Dict]):
    """Print a summary table across all tickers per threshold."""
    print("\n" + "=" * 120)
    print("HYBRID BACKTEST: MACD+RSI → NN QUALITY THRESHOLD SENSITIVITY (OOS 2024-2025)")
    print("=" * 120)

    # Per-threshold summary
    for thr in THRESHOLDS:
        thr_key = str(thr)
        print(f"\n--- nn_entry_threshold = {thr} ---")
        print(f"{'Ticker':<8} {'Model':<20} {'Signals':<8} {'Accepted':<10} {'Trades':<8} {'WR%':<8} {'PF':<8} {'PnL':<12} {'Sharpe':<8} {'MaxDD':<8}")
        print("-" * 100)

        totals = {"accepted": 0, "trades": 0, "pnl": 0, "positive": 0}

        for r in all_results:
            if "error" in r:
                continue
            thr_data = r["thresholds"].get(thr_key, {})
            if not thr_data or thr_data.get("trades", 0) == 0:
                continue

            print(
                f"{r['ticker']:<8} {r.get('model', '?'):<20} "
                f"{thr_data['signals']:<8} {thr_data['accepted']:<10} "
                f"{thr_data['trades']:<8} {thr_data['win_rate']:<8} "
                f"{thr_data['profit_factor']:<8} {thr_data['total_pnl']:<12,.0f} "
                f"{thr_data['sharpe']:<8} {thr_data['max_drawdown']:<8,.0f}"
            )
            totals["accepted"] += thr_data["accepted"]
            totals["trades"] += thr_data["trades"]
            totals["pnl"] += thr_data.get("total_pnl", 0)
            if thr_data.get("total_pnl", 0) > 0:
                totals["positive"] += 1

        print("-" * 100)
        print(
            f"{'TOTAL':<8} {'':<20} {totals['accepted']:<18} "
            f"{totals['trades']:<8} {'':<8} {'':<8} "
            f"{totals['pnl']:<12,.0f} {'':<8} {'':<8}"
        )
        positive_tickers = totals["positive"]
        total_tickers = len([r for r in all_results if "error" not in r])
        print(f"Positive tickers: {positive_tickers}/{total_tickers}")


def main():
    start = "2024-01-01"
    end = "2025-01-01"

    logger.remove()
    logger.add(sys.stdout, level="INFO",
               format="<green>{time:HH:mm:ss}</green> | <level>{message}</level>")

    logger.info(f"HYBRID BACKTEST: {start} → {end}")
    logger.info(f"Tickers: {len(MOEX_TICKERS)}")
    logger.info(f"Thresholds: {THRESHOLDS}")

    all_results = []

    for ticker in MOEX_TICKERS:
        logger.info(f"\n{'='*60}")
        logger.info(f"  {ticker}")
        logger.info(f"{'='*60}")

        result = run_hybrid_backtest(ticker, start, end, THRESHOLDS)

        if "error" not in result:
            best_thr = max(
                result["thresholds"].items(),
                key=lambda x: x[1].get("total_pnl", -1e9)
            )
            logger.info(
                f"  Best threshold: {best_thr[0]} "
                f"(PnL={best_thr[1]['total_pnl']:,.0f}, "
                f"Trades={best_thr[1]['trades']})"
            )
        else:
            logger.warning(f"  ERROR: {result['error']}")

        all_results.append(result)

    # Print summary
    print_summary(all_results)

    # Find best overall threshold
    print(f"\n{'='*120}")
    print("BEST THRESHOLD PER TICKER")
    print(f"{'='*120}")
    print(f"{'Ticker':<8} {'Best Thr':<10} {'PnL':<12} {'WR%':<8} {'PF':<8} {'Trades':<8}")
    print("-" * 60)

    for r in all_results:
        if "error" in r:
            continue
        best_thr, best_data = max(
            r["thresholds"].items(),
            key=lambda x: x[1].get("total_pnl", -1e9)
        )
        if best_data["trades"] > 0:
            print(
                f"{r['ticker']:<8} {best_thr:<10} {best_data['total_pnl']:<12,.0f} "
                f"{best_data['win_rate']:<8} {best_data['profit_factor']:<8} "
                f"{best_data['trades']:<8}"
            )

    # Save results
    RESULTS_DIR.mkdir(exist_ok=True)
    ts = datetime.now().strftime("%Y%m%d_%H%M%S")
    output_path = RESULTS_DIR / f"hybrid_backtest_{ts}.json"
    with open(output_path, "w") as f:
        json.dump(all_results, f, indent=2, default=str)
    logger.info(f"\nResults saved to {output_path}")


if __name__ == "__main__":
    main()
