#!/usr/bin/env python3
"""
Benchmark **implemented** SPY option sleeves on **ThetaData** (default) or **SyntheticLoader**
(``--loader synthetic``) using the same engine as ``vrp_backtest_theta.py``.

Each row runs :class:`~RenTech.strategy_stack.vrp_backtester.VRPBacktester` with a different
``portfolio_mode``: isolated structures in their VIX bands, plus the full production ladder.

This does **not** add new structures (iron condor, BXM, ATM put-write index, 0DTE, etc.). Those
need separate implementations; the table footer lists them as future work.

Usage::

    python RenTech/strategy_stack/theta_strategy_benchmark.py --max-days 400
    python RenTech/strategy_stack/theta_strategy_benchmark.py --loader synthetic --max-days 800
    python RenTech/strategy_stack/theta_strategy_benchmark.py --output-csv /tmp/theta_bench.csv
"""

from __future__ import annotations

import argparse
import csv
import math
import statistics
import sys
from pathlib import Path
from typing import Any, cast

_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.synthetic_data_loader import SyntheticLoader
from RenTech.core.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.vrp_backtester import (
    DEFAULT_STARTING_CAPITAL,
    DEFAULT_YF_END,
    DEFAULT_YF_START,
    PortfolioMode,
    VRPBacktester,
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)
from RenTech.strategy_stack.vrp_strategy_config import (
    DEFAULT_STRATEGY_CONFIG_PATH,
    apply_strategy_params_to_vrp_backtester_module,
    load_strategy_config_file,
)

_DEFAULT_THETA_DIR = _REPO_ROOT / "RenTech" / "data" / "theta_chunks"

BENCHMARK_MODES: list[tuple[PortfolioMode, str, str]] = [
    ("full", "Full stack (R1 + R2a + R2b + R3 + R4)", "all bands"),
    ("r1_strangle", "R1 weekly long strangle", "VIX < 12"),
    ("r2_diagonal", "R2a put diagonal only", "12 ≤ VIX ≤ 20"),
    ("r2_spread", "R2b put credit spread only", "12 ≤ VIX ≤ 20"),
    ("r2_pair", "R2a + R2b only (no R1/R3/R4)", "12 ≤ VIX ≤ 20"),
    ("r3_put_spread", "R3 put credit spread", "20 < VIX ≤ 30"),
    ("r4_credit_spread", "R4 wider put credit spread", "VIX > 30"),
]


def _equity_sharpe_approx(bt: VRPBacktester) -> float:
    ec = bt._equity_curve
    if len(ec) < 3:
        return float("nan")
    caps = [float(c) for _, c in ec]
    rets: list[float] = []
    for i in range(1, len(caps)):
        if caps[i - 1] > 1e-9:
            rets.append((caps[i] - caps[i - 1]) / caps[i - 1])
    if len(rets) < 2:
        return float("nan")
    sd = statistics.pstdev(rets)
    if sd < 1e-12:
        return float("nan")
    mu = statistics.mean(rets)
    ann = (mu / sd) * math.sqrt(252.0)
    return float(ann) if math.isfinite(ann) else float("nan")


def _run_one(
    loader: ThetaChunksLoader | SyntheticLoader,
    spy_wide: pd.DataFrame,
    days: list[pd.Timestamp],
    cap: float,
    mode: PortfolioMode,
    bt_kw: dict[str, Any],
) -> dict[str, Any]:
    kw = dict(bt_kw)
    kw["portfolio_mode"] = mode
    bt = VRPBacktester(loader, **kw)
    bt.run_backtest(trading_days=days, show_progress=False)
    m = bt.metrics()
    row: dict[str, Any] = {
        "portfolio_mode": mode,
        "ending_capital": m["ending_capital"],
        "total_return": m["total_return"],
        "cagr": m["cagr"],
        "max_drawdown": m["max_drawdown"],
        "total_trades": m["total_trades"],
        "sharpe_approx": _equity_sharpe_approx(bt),
        "pmcc_trades": m["pmcc_trades"],
        "diagonal_trades": m["diagonal_trades"],
        "r2_spread_trades": m["r2_spread_trades"],
        "naked_trades": m["naked_trades"],
        "credit_spread_trades": m["credit_spread_trades"],
    }
    return row


def main() -> None:
    ap = argparse.ArgumentParser(
        description="ThetaData benchmark: full VRP vs isolated sleeves (same sizing JSON as production).",
    )
    ap.add_argument(
        "--loader",
        choices=("theta", "synthetic"),
        default="theta",
        help="theta: ThetaData Parquet chunks (default). synthetic: yfinance + SyntheticLoader (no Theta files).",
    )
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA_DIR)
    ap.add_argument("--capital", type=float, default=DEFAULT_STARTING_CAPITAL)
    ap.add_argument("--max-days", type=int, default=0, help="If >0, cap trading days (smoke test).")
    ap.add_argument("--start", type=str, default="", help="YYYY-MM-DD inclusive lower bound.")
    ap.add_argument("--end", type=str, default="", help="YYYY-MM-DD inclusive upper bound.")
    ap.add_argument("--no-progress", action="store_true", help="Unused (runs silent per mode).")
    ap.add_argument("--slippage", type=float, default=None)
    ap.add_argument("--no-vol-scaling", action="store_true")
    ap.add_argument("--no-r2-crossover-filters", action="store_true")
    ap.add_argument("--dd-risk-scaling", action="store_true")
    ap.add_argument("--dd-enter", type=float, default=0.15)
    ap.add_argument("--dd-exit", type=float, default=0.10)
    ap.add_argument("--dd-mult", type=float, default=0.50)
    ap.add_argument("--overlap-portfolio", action="store_true")
    ap.add_argument("--overlap-slice-contracts", type=int, default=1)
    ap.add_argument("--strategy-config", type=Path, default=None)
    ap.add_argument(
        "--modes",
        type=str,
        default="",
        help="Comma-separated portfolio_mode values (default: all built-in benchmark modes).",
    )
    ap.add_argument("--output-csv", type=Path, default=None)
    args = ap.parse_args()

    if args.loader == "theta":
        theta_dir = args.theta_dir.expanduser()
        if not theta_dir.is_dir():
            print(f"ERROR: --theta-dir is not a directory: {theta_dir}", file=sys.stderr)
            sys.exit(1)
        try:
            d0, d1 = theta_chunks_date_bounds(theta_dir)
        except (OSError, ValueError, FileNotFoundError) as e:
            print(f"ERROR: {e}", file=sys.stderr)
            sys.exit(1)
        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))
        ld = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
        days = trading_days_intersecting_spy(ld, spy_wide.index, d0, d1)
        data_label = f"Theta: {theta_dir}"
        window_label = f"{d0.date()} → {d1.date()}"
    else:
        spy_wide = normalize_spy_df(load_spy_vix_from_yfinance(DEFAULT_YF_START, DEFAULT_YF_END))
        ld = SyntheticLoader(spy_wide)
        days = [pd.Timestamp(ts).normalize() for ts in spy_wide.index]
        d0, d1 = days[0], days[-1]
        data_label = "SyntheticLoader (yfinance default window)"
        window_label = f"{d0.date()} → {d1.date()}"
    if not days:
        print("ERROR: No overlapping Theta vs SPY/VIX days.", file=sys.stderr)
        sys.exit(1)
    if args.start.strip():
        t0 = pd.Timestamp(args.start.strip())
        days = [d for d in days if d >= t0]
    if args.end.strip():
        t1 = pd.Timestamp(args.end.strip())
        days = [d for d in days if d <= t1]
    if int(args.max_days) > 0:
        days = days[: int(args.max_days)]
    if not days:
        print("ERROR: No days left after filters.", file=sys.stderr)
        sys.exit(1)

    cap = float(args.capital)
    cfg_path = (
        args.strategy_config.expanduser()
        if args.strategy_config is not None
        else DEFAULT_STRATEGY_CONFIG_PATH
    )
    bt_kw: dict[str, Any] = {
        "initial_capital": cap,
        "spy_df": spy_wide,
        "vol_risk_scaling": not bool(args.no_vol_scaling),
        "r2_crossover_filters": not bool(args.no_r2_crossover_filters),
        "dd_risk_scaling": bool(args.dd_risk_scaling),
        "dd_scale_enter": float(args.dd_enter),
        "dd_scale_exit": float(args.dd_exit),
        "dd_scale_mult": float(args.dd_mult),
        "overlap_portfolio": bool(args.overlap_portfolio),
        "overlap_slice_contracts": int(args.overlap_slice_contracts),
    }
    if cfg_path.is_file():
        _cfg = load_strategy_config_file(cfg_path)
        apply_strategy_params_to_vrp_backtester_module(_cfg.strategy_params)
        bt_kw["sleeve_risk_fractions"] = cast(dict, dict(_cfg.sleeve_risk_fractions))
        bt_kw["overlay_risk_fractions"] = dict(_cfg.overlay_risk_fractions) if _cfg.overlay_risk_fractions else None
        bt_kw["overlay_risk_cap_frac"] = _cfg.overlay_risk_cap_frac
        bt_kw["total_risk_cap_frac"] = _cfg.total_risk_cap_frac
        cfg_note = str(cfg_path.resolve())
    else:
        cfg_note = f"(no JSON at {cfg_path.resolve()})"

    if args.slippage is not None:
        bt_kw["slippage_factor"] = float(args.slippage)

    if args.modes.strip():
        mode_list = [cast(PortfolioMode, m.strip()) for m in args.modes.split(",") if m.strip()]
    else:
        mode_list = [m for m, _, _ in BENCHMARK_MODES]

    labels = {mode: (desc, band) for mode, desc, band in BENCHMARK_MODES}

    rows_out: list[dict[str, Any]] = []
    print("=" * 100)
    print(" Sleeve benchmark | same rules/sizing as VRPBacktester | SPY > SMA(200) for entries")
    print(f" {data_label} | days={len(days)} {window_label} | capital=${cap:,.0f}")
    print(f" Config: {cfg_note}")
    print(
        f" overlap={args.overlap_portfolio} slice={args.overlap_slice_contracts} | "
        f"vol_scaling={'OFF' if args.no_vol_scaling else 'ON'} | r2_filters={'OFF' if args.no_r2_crossover_filters else 'ON'}"
    )
    print("=" * 100)

    hdr = (
        f"{'mode':<20} {'description':<38} {'VIX':<14} {'end$':>12} {'totRet':>9} {'CAGR':>9} {'maxDD':>9} "
        f"{'Sharpe*':>8} {'trades':>6} {'R1':>4} {'R2a':>4} {'R2b':>4} {'R3':>4} {'R4':>4}"
    )
    print(hdr)
    print("-" * len(hdr))
    for mode in mode_list:
        desc, band = labels.get(mode, (mode, ""))
        r = _run_one(ld, spy_wide, days, cap, mode, bt_kw)
        rows_out.append({**r, "description": desc, "vix_band": band})
        cagr = r["cagr"]
        cg_s = f"{cagr:.2%}" if isinstance(cagr, float) and math.isfinite(cagr) else "n/a"
        sh = r["sharpe_approx"]
        sh_s = f"{sh:.2f}" if isinstance(sh, float) and math.isfinite(sh) else "n/a"
        print(
            f"{mode:<20} {desc:<38} {band:<14} "
            f"{r['ending_capital']:>12,.0f} {r['total_return']:>8.2%} {cg_s:>9} {r['max_drawdown']:>8.2%} "
            f"{sh_s:>8} {int(r['total_trades']):>6} "
            f"{int(r['pmcc_trades']):>4} {int(r['diagonal_trades']):>4} {int(r['r2_spread_trades']):>4} "
            f"{int(r['naked_trades']):>4} {int(r['credit_spread_trades']):>4}"
        )

    print("-" * len(hdr))
    print("*Sharpe_approx: mean/std of step-to-step equity changes × sqrt(252); rough, not trade-frequency-adjusted.")
    print("Not implemented here: iron condor, BXM/covered call, CBOE PUT-style ATM cash-secured put, 0DTE grids, VIX flies.")
    print("=" * 100)

    if args.output_csv is not None:
        p = args.output_csv.expanduser()
        p.parent.mkdir(parents=True, exist_ok=True)
        keys = [
            "portfolio_mode",
            "description",
            "vix_band",
            "ending_capital",
            "total_return",
            "cagr",
            "max_drawdown",
            "sharpe_approx",
            "total_trades",
            "pmcc_trades",
            "diagonal_trades",
            "r2_spread_trades",
            "naked_trades",
            "credit_spread_trades",
        ]
        with p.open("w", newline="", encoding="utf-8") as f:
            w = csv.DictWriter(f, fieldnames=keys)
            w.writeheader()
            for row in rows_out:
                w.writerow({k: row.get(k, "") for k in keys})
        print(f"Wrote {p}")


if __name__ == "__main__":
    main()
