#!/usr/bin/env python3
"""Optimized parameter tuning - generates signals once, then filters by threshold.

Usage:
    python scripts/tune_params_v2.py
    python scripts/tune_params_v2.py --sides LONG
"""
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_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

    # Step 1: Generate signals ONE time at lowest threshold
    logger.info(f"Step 1: Generating all signals at low threshold (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.5,  # lowest threshold to capture all candidates
        min_confidence=0.0,   # capture all candidates
    )

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

    # Convert to DataFrame with all metadata
    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"Total raw signals: {len(signals_df)} ({signals_df['signal_type'].value_counts().to_dict()})")

    # Save to cache
    cache_path = Path("reports/tune_signals_cache.csv")
    signals_df.to_csv(cache_path, index=False)
    logger.info(f"Signals cached to {cache_path}")

    # Step 2: Sweep thresholds
    thresholds = [0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9]
    confidences = [0.0, 0.3, 0.4, 0.5, 0.6]

    results = []

    print(f"\n{'='*110}")
    print(f"Parameter tuning for {ticker} ({start} → {end})")
    print(f"Total raw signals: {len(signals_df)}")
    print(f"{'='*110}")
    print(f"{'Threshold':>10} {'MinConf':>8} {'Signals':>8} {'Trades':>8} {'Win%':>8} {'PF':>8} {'AvgRR':>8} {'TotalPnL':>12} {'MaxDD':>10} {'Sharpe':>8}")
    print(f"{'-'*110}")

    for thresh in thresholds:
        for conf in confidences:
            if conf > thresh:
                continue

            # Filter signals by threshold and confidence
            mask = (
                (signals_df["entry_probability"] >= thresh) &
                (signals_df["ai_confidence"] >= conf)
            )
            filtered = signals_df[mask].copy()

            if filtered.empty:
                print(f"{thresh:>10.2f} {conf:>8.2f} {'0':>8} {'0':>8} {'N/A':>8} {'N/A':>8} {'N/A':>8} {'N/A':>12} {'N/A':>10} {'N/A':>8}")
                results.append({
                    "threshold": thresh,
                    "min_confidence": conf,
                    "signals": 0,
                    "error": "No signals after filter",
                })
                continue

            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"{thresh:>10.2f} {conf:>8.2f} {len(filtered):>8} {'0':>8} {'N/A':>8} {'N/A':>8} {'N/A':>8} {'N/A':>12} {'N/A':>10} {'N/A':>8}")
                results.append({
                    "threshold": thresh,
                    "min_confidence": conf,
                    "signals": len(filtered),
                    "error": "No trades executed",
                })
                continue

            metrics = calculate_metrics(trades, backtester.equity)

            results.append({
                "threshold": thresh,
                "min_confidence": conf,
                "signals": len(filtered),
                "total_trades": metrics["total_trades"],
                "win_rate": metrics["win_rate"],
                "profit_factor": metrics["profit_factor"],
                "avg_rr": metrics["avg_rr"],
                "total_pnl": metrics["total_pnl"],
                "max_drawdown": metrics["max_drawdown"],
                "sharpe_ratio": metrics.get("sharpe_ratio", 0),
                "expectancy": metrics.get("expectancy", 0),
                "avg_win": metrics.get("avg_win", 0),
                "avg_loss": metrics.get("avg_loss", 0),
            })

            print(
                f"{thresh:>10.2f} {conf:>8.2f} "
                f"{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} "
                f"{metrics['max_drawdown']:>10.0f} {metrics.get('sharpe_ratio', 0):>8.2f}"
            )

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

    if results:
        # Best by profit factor (min 30 trades)
        valid = [r for r in results if r.get("total_trades", 0) >= 30]
        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:     thresh={best_pf['threshold']:.2f}, conf={best_pf['min_confidence']:.2f}, "
                  f"PF={best_pf['profit_factor']:.2f}, PnL={best_pf['total_pnl']:,.0f}, "
                  f"WR={best_pf['win_rate']:.1f}%, trades={best_pf['total_trades']}")
            print(f"Best PnL:    thresh={best_pnl['threshold']:.2f}, conf={best_pnl['min_confidence']:.2f}, "
                  f"PF={best_pnl['profit_factor']:.2f}, PnL={best_pnl['total_pnl']:,.0f}, "
                  f"WR={best_pnl['win_rate']:.1f}%, trades={best_pnl['total_trades']}")
            print(f"Best Sharpe: thresh={best_sharpe['threshold']:.2f}, conf={best_sharpe['min_confidence']:.2f}, "
                  f"Sharpe={best_sharpe['sharpe_ratio']:.2f}, PnL={best_sharpe['total_pnl']:,.0f}, "
                  f"WR={best_sharpe['win_rate']:.1f}%, trades={best_sharpe['total_trades']}")

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


if __name__ == "__main__":
    run_tuning()
