#!/usr/bin/env python3
"""Parameter tuning script for neural network backtester.

Tests different ai_threshold and min_confidence combinations on a single ticker
and reports results in a table.
"""
from __future__ import annotations

import argparse
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))

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


def test_params(
    ticker: str,
    start: str,
    end: str,
    model_path: Path,
    ai_threshold: float,
    min_confidence: float,
    sides: List[str],
    capital: float = 1_000_000,
    risk: float = 0.01,
    max_positions: int = 5,
) -> Dict[str, Any]:
    """Run a single backtest with given params and return metrics."""
    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)

    if df_h1.empty:
        return {"error": "No data"}

    predictor = NeuralPredictor(
        model_path=model_path,
        entry_threshold=ai_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:
        return {"error": "No signals"}

    signals_df = predictor.signals_to_dataframe(all_signals)

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

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

    metrics = calculate_metrics(trades, backtester.equity)
    metrics["signals"] = len(signals_df)
    metrics["trades"] = len(trades)
    return metrics


def main():
    parser = argparse.ArgumentParser(description="Parameter tuning")
    parser.add_argument("--ticker", type=str, default="SBER")
    parser.add_argument("--start", type=str, default="2023-01-01")
    parser.add_argument("--end", type=str, default="2024-01-01")
    parser.add_argument("--ai-model", type=str, default="models/neural_trader.pt")
    parser.add_argument("--capital", type=float, default=1_000_000)
    parser.add_argument("--risk", type=float, default=0.01)
    parser.add_argument("--sides", type=str, default="LONG,SHORT")
    args = parser.parse_args()

    model_path = Path(args.ai_model)
    sides = args.sides.split(",")

    # Threshold sweep values
    thresholds = [0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9]
    confidences = [0.3, 0.4, 0.5, 0.6]

    print(f"\nParameter tuning for {args.ticker} ({args.start} → {args.end})")
    print(f"Model: {model_path}, Sides: {sides}")
    print(f"\n{'Threshold':>10} {'MinConf':>8} {'Signals':>8} {'Trades':>8} {'Win%':>8} {'PF':>8} {'AvgRR':>8} {'TotalPnL':>12} {'MaxDD':>10} {'Expect':>8}")
    print("-" * 100)

    results = []
    for thresh in thresholds:
        for conf in confidences:
            if conf > thresh:
                continue  # confidence can't be higher than threshold

            try:
                metrics = test_params(
                    ticker=args.ticker,
                    start=args.start,
                    end=args.end,
                    model_path=model_path,
                    ai_threshold=thresh,
                    min_confidence=conf,
                    sides=sides,
                    capital=args.capital,
                    risk=args.risk,
                )

                if "error" in metrics:
                    print(f"{thresh:>10.2f} {conf:>8.2f} {'ERR':>16} {metrics['error']}")
                    continue

                results.append({
                    "threshold": thresh,
                    "min_confidence": conf,
                    **metrics,
                })

                print(
                    f"{thresh:>10.2f} {conf:>8.2f} "
                    f"{metrics.get('signals',0):>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('expectancy',0):>8.0f}"
                )

            except Exception as e:
                print(f"{thresh:>10.2f} {conf:>8.2f} {'ERR':>16} {e}")

    if results:
        print("\n" + "=" * 100)

        # 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"])
            print(f"\nBest PF:  threshold={best_pf['threshold']}, min_conf={best_pf['min_confidence']}, "
                  f"PF={best_pf['profit_factor']:.2f}, PnL={best_pf['total_pnl']:,.0f}")
            print(f"Best PnL: threshold={best_pnl['threshold']}, min_conf={best_pnl['min_confidence']}, "
                  f"PF={best_pnl['profit_factor']:.2f}, PnL={best_pnl['total_pnl']:,.0f}")

        # Save to file
        out_path = Path(__file__).parent.parent / "reports" / f"tune_{args.ticker}_{args.start}_{args.end}.json"
        with open(out_path, "w") as f:
            # Only keep summary fields
            summary = [{
                "threshold": r["threshold"],
                "min_confidence": r["min_confidence"],
                "signals": r.get("signals", 0),
                "trades": r["total_trades"],
                "win_rate": r["win_rate"],
                "profit_factor": r["profit_factor"],
                "avg_rr": r["avg_rr"],
                "total_pnl": r["total_pnl"],
                "max_drawdown": r["max_drawdown"],
                "expectancy": r.get("expectancy", 0),
                "sharpe_ratio": r.get("sharpe_ratio", 0),
            } for r in results]
            json.dump(summary, f, indent=2)
        print(f"\nResults saved to {out_path}")


if __name__ == "__main__":
    main()
