#!/usr/bin/env python3
"""Tune TP multiplier (RR) on top of best threshold config.

Overrides the model's TP with: entry ± (entry - sl) * tp_rr_mult
Then runs batch backtest with best threshold + various RR multipliers.
"""
from __future__ import annotations

import json
import sys
from datetime import datetime
from pathlib import Path
from typing import List, Dict, Any

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

import pandas as pd
from loguru import logger

from core.data_loader import DataPreparator
from ai.inference_v2 import NeuralPredictor
from backtest.engine import NeuralBacktester
from backtest.metrics import calculate_metrics


def run_tp_tuning():
    ticker = "SBER"
    start = "2023-01-01"
    end = "2024-01-01"
    model_path = Path("models/neural_trader.pt")
    sides = ["LONG", "SHORT"]
    capital = 1_000_000
    risk_pct = 0.01
    max_positions = 5

    # Generate signals once at best threshold (0.9, conf=0.5)
    logger.info("Generating signals at threshold=0.9, min_conf=0.5")
    start_ts = int(datetime.fromisoformat(start).timestamp())
    end_ts = int(datetime.fromisoformat(end).timestamp())

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

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

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

    # Convert to DataFrame
    rows = []
    for s in all_signals:
        rows.append({
            "timestamp": s.timestamp,
            "signal_type": s.side,
            "entry_price": s.entry_price,
            "sl_price": s.sl_price,
            "tp_price": s.tp_price,
            "sl_distance_atr": s.sl_distance_atr,
            "tp_distance_atr": s.tp_distance_atr,
            "entry_probability": s.entry_probability,
            "ai_confidence": s.confidence,
            "atr_at_signal": s.atr_at_signal,
        })

    signals_df = pd.DataFrame(rows)
    signals_df = signals_df.sort_values("timestamp").reset_index(drop=True)
    logger.info(f"Signals: {len(signals_df)}")

    # Test various TP multipliers
    rr_multipliers = [1.0, 1.5, 2.0, 2.5, 3.0, 4.0]

    print(f"\n{'='*100}")
    print(f"TP/RR Tuning for {ticker} ({start} → {end})")
    print(f"Threshold=0.9, MinConf=0.5, Signals={len(signals_df)}")
    print(f"{'='*100}")
    print(f"{'TP_RR':>8} {'Signals':>8} {'Trades':>8} {'Win%':>8} {'PF':>8} {'AvgRR':>8} {'TotalPnL':>12} {'MaxDD':>10} {'Sharpe':>8} {'AvgWin':>8} {'AvgLoss':>8}")
    print(f"{'-'*100}")

    results = []
    for mult in rr_multipliers:
        # Override TP in signals
        df = signals_df.copy()
        for idx, row in df.iterrows():
            entry = row["entry_price"]
            sl = row["sl_price"]
            side = row["signal_type"]
            risk = abs(entry - sl)
            if side == "LONG":
                df.at[idx, "tp_price"] = entry + risk * mult
            else:
                df.at[idx, "tp_price"] = entry - risk * mult

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

        if not trades:
            print(f"{mult:>8.1f} {len(df):>8} {'0':>8} {'N/A':>8} {'N/A':>8} {'N/A':>8} {'N/A':>12} {'N/A':>10} {'N/A':>8} {'N/A':>8} {'N/A':>8}")
            continue

        metrics = calculate_metrics(trades, backtester.equity)
        results.append({
            "tp_rr_mult": mult,
            "signals": len(df),
            **metrics,
        })

        print(
            f"{mult:>8.1f} {len(df):>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} "
            f"{metrics['max_drawdown']:>10.0f} {metrics.get('sharpe_ratio',0):>8.2f} "
            f"{metrics.get('avg_win',0):>8.0f} {metrics.get('avg_loss',0):>8.0f}"
        )

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

    if results:
        # Best by various metrics
        valid = [r for r in results if r.get("total_trades", 0) >= 10]
        if valid:
            best_pf = max(valid, key=lambda x: x["profit_factor"])
            best_pnl = max(valid, key=lambda x: x["total_pnl"])
            best_sharpe = max(valid, key=lambda x: x.get("sharpe_ratio", -999))

            print(f"\nBest PF:     mult={best_pf['tp_rr_mult']:.1f}, PF={best_pf['profit_factor']:.2f}, "
                  f"PnL={best_pf['total_pnl']:,.0f}, WR={best_pf['win_rate']:.1f}%")
            print(f"Best PnL:    mult={best_pnl['tp_rr_mult']:.1f}, PF={best_pnl['profit_factor']:.2f}, "
                  f"PnL={best_pnl['total_pnl']:,.0f}, WR={best_pnl['win_rate']:.1f}%")
            print(f"Best Sharpe: mult={best_sharpe['tp_rr_mult']:.1f}, Sharpe={best_sharpe['sharpe_ratio']:.2f}, "
                  f"PnL={best_sharpe['total_pnl']:,.0f}")

        # Save
        out_path = Path("reports") / f"tune_rr_{ticker}_{start}_{end}.json"
        with open(out_path, "w") as f:
            json.dump(results, f, indent=2)
        print(f"\nSaved to {out_path}")


if __name__ == "__main__":
    run_tp_tuning()
