#!/usr/bin/env python3
"""
Sweep R2a (``diagonal`` sleeve) vs R2b (``r2_spread`` sleeve) risk fractions on the
Theta VRP backtest. Loads chains once, reuses ``trading_days`` for each run.

Example::

    python RenTech/strategy_stack/sweep_r2_sleeve_fracs.py \\
        --start 2022-01-03 --end 2024-12-31 --capital 100000

Defaults use ``sleeve_risk_fractions.json`` for non-R2 sleeves and strategy_params.
"""

from __future__ import annotations

import argparse
import json
import math
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.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.vrp_backtester import VRPBacktester
from RenTech.strategy_stack.vrp_backtest_theta import (
    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"


def _parse_grid(s: str) -> list[tuple[float, float]]:
    """Comma-separated rows ``r2a,r2b;r2a,r2b`` or ``r2a:r2b`` pairs."""
    rows: list[tuple[float, float]] = []
    for chunk in s.replace(";", ",").split(","):
        chunk = chunk.strip()
        if not chunk:
            continue
        if ":" in chunk:
            a, b = chunk.split(":", 1)
            rows.append((float(a), float(b)))
        else:
            raise ValueError(f"Bad grid token {chunk!r}; use r2a:r2b pairs separated by commas")
    return rows


def main() -> None:
    ap = argparse.ArgumentParser(description="Sweep R2a/R2b sleeve risk fractions (VRP Theta backtest).")
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA_DIR)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--start", type=str, default="", help="YYYY-MM-DD inclusive lower bound on trading days.")
    ap.add_argument("--end", type=str, default="", help="YYYY-MM-DD inclusive upper bound.")
    ap.add_argument("--max-days", type=int, default=0, help="If >0, only first N days after filters (debug).")
    ap.add_argument(
        "--strategy-config",
        type=Path,
        default=None,
        help="Parity JSON (default: RenTech/strategy_stack/sleeve_risk_fractions.json).",
    )
    ap.add_argument(
        "--grid",
        type=str,
        default="",
        help="Override default grid: comma-separated ``r2a:r2b`` pairs, e.g. ``0.02:0.01,0.04:0.02``.",
    )
    ap.add_argument("--no-vol-scaling", action="store_true")
    ap.add_argument("--no-r2-crossover-filters", action="store_true")
    ap.add_argument(
        "--enforce-total-risk-cap",
        action="store_true",
        help="Use total_risk_cap_frac from JSON (default: off so R2a>R1 can exceed 3%% cap during sweeps).",
    )
    args = ap.parse_args()

    cfg_path = (
        args.strategy_config.expanduser()
        if args.strategy_config is not None
        else DEFAULT_STRATEGY_CONFIG_PATH
    )
    if not cfg_path.is_file():
        print(f"ERROR: strategy config not found: {cfg_path}", file=sys.stderr)
        sys.exit(1)

    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)

    _cfg = load_strategy_config_file(cfg_path)
    apply_strategy_params_to_vrp_backtester_module(_cfg.strategy_params)

    base_sleeves = dict(_cfg.sleeve_risk_fractions)
    if args.grid.strip():
        grid = _parse_grid(args.grid.strip())
    else:
        grid = [
            (0.020, 0.010),
            (0.025, 0.0125),
            (0.030, 0.015),
            (0.035, 0.0175),
            (0.040, 0.020),
            (0.045, 0.0225),
            (0.050, 0.025),
            (0.060, 0.030),
            (0.040, 0.010),
            (0.020, 0.020),
            (0.050, 0.020),
        ]

    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))
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_wide)
    days = trading_days_intersecting_spy(ld, spy_wide.index, d0, d1)
    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 trading days after filters.", file=sys.stderr)
        sys.exit(1)

    cap = float(args.capital)
    bt_base: 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": False,
        "overlay_risk_fractions": dict(_cfg.overlay_risk_fractions) if _cfg.overlay_risk_fractions else None,
        "overlay_risk_cap_frac": _cfg.overlay_risk_cap_frac,
        "total_risk_cap_frac": _cfg.total_risk_cap_frac if args.enforce_total_risk_cap else None,
    }

    print(
        json.dumps(
            {
                "theta_dir": str(theta_dir.resolve()),
                "strategy_config": str(cfg_path.resolve()),
                "n_days": len(days),
                "first_day": str(days[0].date()),
                "last_day": str(days[-1].date()),
                "capital": cap,
                "grid_runs": len(grid),
            },
            indent=2,
        ),
        flush=True,
    )
    span_years = max((days[-1] - days[0]).days / 365.25, 1e-9)
    print(
        f"# backtest_span_years={span_years:.4f} (CAGR_win uses this; engine CAGR uses full SPY yfinance span)",
        flush=True,
    )
    print(
        "r2a_frac\tr2b_frac\tend_cap\ttot_ret\tmax_dd\tcagr_win\tcagr_engine\tn_trades",
        flush=True,
    )
    results: list[dict[str, Any]] = []
    for r2a, r2b in grid:
        sleeves = dict(base_sleeves)
        sleeves["diagonal"] = float(r2a)
        sleeves["r2_spread"] = float(r2b)
        bt = VRPBacktester(ld, **bt_base, sleeve_risk_fractions=cast(dict, sleeves))
        bt.run_backtest(trading_days=days, show_progress=False)
        m = bt.metrics()
        cagr_eng = m["cagr"]
        endv = float(m["ending_capital"])
        cagr_win = (endv / cap) ** (1.0 / span_years) - 1.0 if cap > 0 and endv > 0 else float("nan")
        ce_s = f"{float(cagr_eng):.4f}" if isinstance(cagr_eng, float) and math.isfinite(cagr_eng) else "nan"
        cw_s = f"{cagr_win:.4f}" if math.isfinite(cagr_win) else "nan"
        print(
            f"{r2a:.4f}\t{r2b:.4f}\t{m['ending_capital']:.2f}\t{m['total_return']:.4f}\t"
            f"{m['max_drawdown']:.4f}\t{cw_s}\t{ce_s}\t{m['total_trades']}",
            flush=True,
        )
        results.append(
            {
                "r2a_diagonal_frac": r2a,
                "r2b_spread_frac": r2b,
                "ending_capital": m["ending_capital"],
                "total_return": m["total_return"],
                "max_drawdown": m["max_drawdown"],
                "cagr_window": cagr_win,
                "cagr_engine": cagr_eng,
                "total_trades": m["total_trades"],
            }
        )

    # Rank by closeness to target CAGR 14% and max DD ~10% (simple L2 on window CAGR + reported max DD).
    target_c, target_dd = 0.14, 0.10

    def score(row: dict[str, Any]) -> float:
        c = row["cagr_window"]
        dd = row["max_drawdown"]
        if not (isinstance(c, float) and math.isfinite(c)):
            return 1e9
        return (c - target_c) ** 2 + (dd - target_dd) ** 2

    best = min(results, key=score)
    print(
        "--- closest to targets (CAGR_win≈14%, maxDD≈10%) under L2 ---\n"
        f"best_r2a={best['r2a_diagonal_frac']:.4f} best_r2b={best['r2b_spread_frac']:.4f} "
        f"cagr_window={best['cagr_window']:.4f} max_dd={best['max_drawdown']:.4f}",
        flush=True,
    )


if __name__ == "__main__":
    main()
