#!/usr/bin/env python3
"""Quick test: thr=0.85 + tp_rr=3.0 on SBER, then thr=0.80

Uses cached signals from tune_params_v2 run.
"""
from __future__ import annotations

import json
import sys
from pathlib import Path

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

import pandas as pd
from loguru import logger

from core.data_loader import DataPreparator
from backtest.engine import NeuralBacktester
from backtest.metrics import calculate_metrics


def run():
    ticker = "SBER"
    capital = 1_000_000
    risk_pct = 0.01
    max_positions = 5

    # Load cached signals
    cache_path = Path("reports/tune_signals_cache.csv")
    if not cache_path.exists():
        logger.error("Cache not found! Run tune_params_v2.py first")
        return

    signals_df = pd.read_csv(cache_path)
    logger.info(f"Loaded {len(signals_df)} cached signals")

    # Need OHLCV for backtest (use same period)
    prep = DataPreparator([ticker])
    from datetime import datetime
    start_ts = int(datetime(2023,1,1).timestamp())
    end_ts = int(datetime(2024,1,1).timestamp())
    df_h1 = prep.load_h1(ticker, start_ts, end_ts)

    # Test configurations
    configs = [
        {"thr": 0.90, "tp_rr": 3.0, "label": "thr=0.9, RR=3 (baseline)"},
        {"thr": 0.85, "tp_rr": 3.0, "label": "thr=0.85, RR=3"},
        {"thr": 0.80, "tp_rr": 3.0, "label": "thr=0.80, RR=3"},
        {"thr": 0.75, "tp_rr": 3.0, "label": "thr=0.75, RR=3"},
    ]

    print(f"\n{'='*90}")
    print(f"Trade count vs threshold on {ticker}")
    print(f"{'='*90}")
    print(f"{'Config':<28} {'Signals':>8} {'Trades':>8} {'Win%':>8} {'PF':>8} {'AvgRR':>8} {'TotalPnL':>12}")
    print(f"{'-'*90}")

    for cfg in configs:
        thr = cfg["thr"]
        tp_rr = cfg["tp_rr"]

        # Filter signals by threshold (confidence already baked in)
        mask = signals_df["entry_probability"] >= thr
        filtered = signals_df[mask].copy()

        if filtered.empty:
            print(f"{cfg['label']:<28} {0:>8} {0:>8} {'N/A':>8} {'N/A':>8} {'N/A':>8} {'N/A':>12}")
            continue

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

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

        if not trades:
            print(f"{cfg['label']:<28} {len(filtered):>8} {0:>8} {'N/A':>8} {'N/A':>8} {'N/A':>8} {'N/A':>12}")
            continue

        metrics = calculate_metrics(trades, backtester.equity)

        print(
            f"{cfg['label']:<28} {len(filtered):>8} {metrics['total_trades']:>8} "
            f"{metrics['win_rate']:>8.1f} {metrics['profit_factor']:>8.2f} "
            f"{metrics['avg_rr']:>8.2f} {metrics['total_pnl']:>12.0f}"
        )

    print(f"{'-'*90}")


if __name__ == "__main__":
    run()
