#!/usr/bin/env python3
"""Batch backtester — runs neural network backtest on all tickers.

Usage:
    python scripts/batch_backtest.py
    python scripts/batch_backtest.py --ticker SBER,GAZP
    python scripts/batch_backtest.py --start 2023-01-01 --end 2024-01-01
    python scripts/batch_backtest.py --ai-model models/neural_trader.pt
"""

from __future__ import annotations

import argparse
import json
import sys
from datetime import datetime
from pathlib import Path

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.inference_v2 import NeuralPredictor
from backtest.engine import NeuralBacktester
from backtest.metrics import calculate_metrics, save_full_report


MOEX_TICKERS = [
    "ASTR", "GAZP", "LKOH", "MTSS", "NVTK",
    "PHOR", "PLZL", "ROSN", "SBER", "SNGSP",
    "VTBR", "X5",
]

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


def run_single_backtest(
    ticker: str,
    start: str,
    end: str,
    sides: list,
    model_path: Path,
    entry_threshold: float = 0.6,
    min_confidence: float = 0.3,
    capital: float = 1_000_000,
    risk: float = 0.01,
    max_positions: int = 5,
    tp_rr_mult: float = 0.0,
) -> dict:
    """Run a single neural network backtest and return metrics dict."""
    start_ts = int(datetime.fromisoformat(start).timestamp())
    end_ts = int(datetime.fromisoformat(end).timestamp())

    try:
        prep = DataPreparator([ticker])
        df_h1 = prep.load_h1(ticker, start_ts, end_ts)

        if df_h1.empty:
            logger.warning(f"No H1 data for {ticker}")
            return {"error": "No H1 data", "ticker": ticker, "sides": sides}

        if not model_path.exists():
            return {"error": f"Model not found: {model_path}", "ticker": ticker}

        predictor = NeuralPredictor(
            model_path=model_path,
            entry_threshold=entry_threshold,
            min_confidence=min_confidence,
        )

        all_signals = []
        for side in sides:
            signals = predictor.predict_signals(df_h1, side=side)
            all_signals.extend(signals)

        if not all_signals:
            logger.warning(f"No signals for {ticker}")
            return {"error": "No signals", "ticker": ticker, "sides": sides}

        signals_df = predictor.signals_to_dataframe(all_signals)

        # Override TP with RR multiplier if configured
        if tp_rr_mult > 0:
            for idx, row in signals_df.iterrows():
                entry = row["entry_price"]
                sl = row["sl_price"]
                side = row["signal_type"]
                risk = abs(entry - sl)
                if side == "LONG":
                    signals_df.at[idx, "tp_price"] = entry + risk * tp_rr_mult
                else:
                    signals_df.at[idx, "tp_price"] = entry - risk * tp_rr_mult

        backtester = NeuralBacktester(
            capital=capital, risk_pct=risk, max_positions=max_positions,
        )
        trades = backtester.run(signals_df, df_h1, ticker=ticker)

        if not trades:
            return {
                "ticker": ticker, "sides": sides,
                "signals": len(signals_df), "trades": 0,
                "error": "No trades executed",
            }

        metrics = calculate_metrics(trades, backtester.equity)

        return {
            "ticker": ticker,
            "sides": sides,
            "signals": len(signals_df),
            **metrics,
        }

    except Exception as e:
        logger.error(f"Backtest failed for {ticker}: {e}")
        return {"error": str(e), "ticker": ticker, "sides": sides}


def main():
    parser = argparse.ArgumentParser(description="Batch neural network backtester")
    parser.add_argument("--ticker", type=str, default=None, help="Comma-separated tickers")
    parser.add_argument("--start", type=str, default="2023-01-01")
    parser.add_argument("--end", type=str, default="2024-01-01")
    parser.add_argument("--capital", type=float, default=1_000_000)
    parser.add_argument("--risk", type=float, default=0.01)
    parser.add_argument("--ai-model", type=str, default="models/neural_trader.pt")
    parser.add_argument("--ai-threshold", type=float, default=0.6)
    parser.add_argument("--min-confidence", type=float, default=0.3)
    parser.add_argument("--sides", type=str, default="LONG,SHORT")
    parser.add_argument("--max-positions", type=int, default=5)
    parser.add_argument("--tp-rr-mult", type=float, default=0.0,
                        help="Override TP: entry ± (entry - sl) * mult. 0 = use model TP")
    args = parser.parse_args()

    tickers = args.ticker.split(",") if args.ticker else MOEX_TICKERS
    sides = args.sides.split(",")
    model_path = Path(args.ai_model)

    logger.info(f"Batch backtest: {len(tickers)} tickers, {args.start} → {args.end}")
    logger.info(f"Model: {model_path}, sides: {sides}, threshold: {args.ai_threshold}")

    if not model_path.exists():
        logger.error(f"Model not found: {model_path}")
        logger.info("Train model first: python scripts/train_neural_model.py")
        return

    all_results = []

    for ticker in tickers:
        logger.info(f"\n{'='*60}\n  Running {ticker}\n{'='*60}")
        result = run_single_backtest(
            ticker, args.start, args.end,
            sides=sides,
            model_path=model_path,
            entry_threshold=args.ai_threshold,
            min_confidence=args.min_confidence,
            capital=args.capital,
            risk=args.risk,
            max_positions=args.max_positions,
            tp_rr_mult=args.tp_rr_mult,
        )
        all_results.append(result)
        if "error" not in result:
            logger.info(
                f"  {ticker}: trades={result.get('total_trades',0)}, "
                f"WR={result.get('win_rate',0):.1f}%, "
                f"PF={result.get('profit_factor',0):.2f}, "
                f"PnL={result.get('total_pnl',0):,.0f}"
            )

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

    print_summary(all_results)


def print_summary(results: list):
    """Print a formatted summary table."""
    print("\n" + "=" * 100)
    print("BATCH BACKTEST SUMMARY (Neural Network)")
    print("=" * 100)
    print(f"{'Ticker':<8} {'Trades':<8} {'Win%':<8} {'PF':<8} {'MaxDD':<10} {'AvgRR':<8} {'TotalPnL':<12} {'Expectancy':<10}")
    print("-" * 100)

    for r in results:
        if "error" in r:
            print(f"{r['ticker']:<8} {'ERR':<8} {r['error'][:50]}")
            continue

        print(
            f"{r['ticker']:<8} {r['total_trades']:<8} "
            f"{r['win_rate']:<8.1f} {r['profit_factor']:<8.2f} "
            f"{r['max_drawdown']:<10.0f} {r['avg_rr']:<8.2f} "
            f"{r['total_pnl']:<12.0f} {r.get('expectancy',0):<10.0f}"
        )

    print("=" * 100)

    valid = [r for r in results if "error" not in r and r.get("total_trades", 0) > 0]
    if valid:
        best = max(valid, key=lambda x: x["total_pnl"])
        worst = min(valid, key=lambda x: x["total_pnl"])
        print(f"\nBest:  {best['ticker']} - PnL={best['total_pnl']:,.0f}, WR={best['win_rate']:.1f}%")
        print(f"Worst: {worst['ticker']} - PnL={worst['total_pnl']:,.0f}, WR={worst['win_rate']:.1f}%")

        total_pnl = sum(r["total_pnl"] for r in valid)
        avg_pf = sum(r["profit_factor"] for r in valid) / len(valid)
        avg_wr = sum(r["win_rate"] for r in valid) / len(valid)
        print(f"\nTotal PnL: {total_pnl:,.0f}, Avg PF: {avg_pf:.2f}, Avg WR: {avg_wr:.1f}%")


if __name__ == "__main__":
    main()
