#!/usr/bin/env python3
"""Compare VRP backtest with vs without VVIX/VIX risk scaling (same Theta window)."""

from __future__ import annotations

import argparse
import math
import sys
from pathlib import Path

_REPO_ROOT = Path(__file__).resolve().parents[2]
if str(_REPO_ROOT) not in sys.path:
    sys.path.insert(0, str(_REPO_ROOT))

import pandas as pd

from RenTech.core.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.vrp_backtester import (
    DEFAULT_STARTING_CAPITAL,
    VRPBacktester,
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)

_DEFAULT_THETA = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"


def _run(
    theta_dir: Path,
    days: list[pd.Timestamp],
    spy_wide: pd.DataFrame,
    vol_risk_scaling: bool,
) -> dict[str, float]:
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
    bt = VRPBacktester(
        ld,
        initial_capital=DEFAULT_STARTING_CAPITAL,
        spy_df=spy_wide,
        vol_risk_scaling=vol_risk_scaling,
        r2_crossover_filters=True,
    )
    bt.run_backtest(trading_days=days, show_progress=False)
    m = bt.metrics()
    return {
        "max_dd": float(m["max_drawdown"]),
        "total_return": float(m["total_return"]),
        "cagr": float(m["cagr"]) if math.isfinite(float(m["cagr"])) else float("nan"),
        "trades": float(m["total_trades"]),
        "end": float(m["ending_capital"]),
        "r2_trades": float(m["diagonal_trades"]),
    }


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA)
    ap.add_argument("--start", type=str, default="2016-01-04")
    ap.add_argument("--end", type=str, default="2026-04-02")
    ap.add_argument("--max-days", type=int, default=0)
    args = ap.parse_args()

    theta_dir = args.theta_dir.expanduser()
    d0, d1 = theta_chunks_date_bounds(theta_dir)
    yf_start = (d0 - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    yf_end = (d1 + pd.Timedelta(days=14)).strftime("%Y-%m-%d")
    spy_wide = normalize_spy_df(load_spy_vix_from_yfinance(yf_start, yf_end))
    ld0 = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
    days = trading_days_intersecting_spy(ld0, spy_wide.index, d0, d1)
    if args.start.strip():
        days = [d for d in days if d >= pd.Timestamp(args.start.strip())]
    if args.end.strip():
        days = [d for d in days if d <= pd.Timestamp(args.end.strip())]
    if int(args.max_days) > 0:
        days = days[: int(args.max_days)]

    print(f"Days: {len(days)}  |  VVIX column: {'vvix_close' in spy_wide.columns}")
    print()

    rows = []
    for label, vs in (("baseline (no vol scaling)", False), ("VVIX + R2 stress scaling", True)):
        m = _run(theta_dir, days, spy_wide, vs)
        rows.append((label, m))
        print(label)
        print(f"  max DD: {m['max_dd']:.2%}  |  return: {m['total_return']:.2%}  |  CAGR: {m['cagr']:.2%}  |  trades: {int(m['trades'])}  |  R2: {int(m['r2_trades'])}")

    b, s = rows[0][1], rows[1][1]
    print()
    print("Delta (scaled − baseline):")
    print(f"  max DD: {(s['max_dd'] - b['max_dd']) * 100:+.2f} pp")
    print(f"  return: {(s['total_return'] - b['total_return']) * 100:+.2f} pp")


if __name__ == "__main__":
    main()
