#!/usr/bin/env python3
"""
Sweep VRPBacktester runtime toggles (same wiring as vrp_backtest_theta.py) and write:

  - RenTech/strategy_stack/VRP_TOGGLE_MATRIX.json
  - RenTech/strategy_stack/VRP_TOGGLE_MATRIX.md

Re-run with ``--max-days 0`` for the longest available overlap (slow).

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_vrp_toggle_matrix.py \\
        --max-days 756 --no-progress
"""

from __future__ import annotations

import argparse
import json
import sys
import time
from dataclasses import asdict, dataclass
from pathlib import Path
import math
from typing import Any, cast

import pandas as pd

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

from RenTech.core.theta_chunks_loader import ThetaChunksLoader, theta_chunks_date_bounds
from RenTech.strategy_stack.vrp_backtester import (
    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"
_JSON_OUT = Path(__file__).resolve().parent / "VRP_TOGGLE_MATRIX.json"
_MD_OUT = Path(__file__).resolve().parent / "VRP_TOGGLE_MATRIX.md"


@dataclass(frozen=True)
class Scenario:
    name: str
    vol_scaling: bool
    r2_crossover: bool
    dd_risk_scaling: bool
    overlap_portfolio: bool
    macro_overlay: bool


def _scenario_grid(include_macro_factorial: bool) -> list[Scenario]:
    """
    Default: 18 scenarios — factorial(vol, r2, dd, overlap) with macro OFF (16 rows), plus
    two explicit macro=Y reference rows (production defaults + “all gates off”).

    With ``include_macro_factorial``: 32 = full 2^5 grid (vol × r2 × dd × ov × macro).
    """
    out: list[Scenario] = []
    mac_levels = (False, True) if include_macro_factorial else (False,)
    for vol in (True, False):
        for r2 in (True, False):
            for dd in (True, False):
                for ov in (True, False):
                    for mac in mac_levels:
                        parts = [
                            f"vol={'Y' if vol else 'n'}",
                            f"r2x={'Y' if r2 else 'n'}",
                            f"dd={'Y' if dd else 'n'}",
                            f"ov={'Y' if ov else 'n'}",
                            f"mac={'Y' if mac else 'n'}",
                        ]
                        out.append(
                            Scenario(
                                name=" | ".join(parts),
                                vol_scaling=vol,
                                r2_crossover=r2,
                                dd_risk_scaling=dd,
                                overlap_portfolio=ov,
                                macro_overlay=mac,
                            )
                        )
    if not include_macro_factorial:
        out.append(
            Scenario(
                name="ref | baseline+macro (vol=Y r2x=Y dd=n ov=n mac=Y)",
                vol_scaling=True,
                r2_crossover=True,
                dd_risk_scaling=False,
                overlap_portfolio=False,
                macro_overlay=True,
            )
        )
        out.append(
            Scenario(
                name="ref | all gates off +macro (vol=n r2x=n dd=n ov=n mac=Y)",
                vol_scaling=False,
                r2_crossover=False,
                dd_risk_scaling=False,
                overlap_portfolio=False,
                macro_overlay=True,
            )
        )
    return out


def _run_one(
    ld: ThetaChunksLoader,
    spy_wide: pd.DataFrame,
    days: list[pd.Timestamp],
    capital: float,
    cfg_path: Path,
    sc: Scenario,
    *,
    dd_enter: float,
    dd_exit: float,
    dd_mult: float,
    macro_frac: float,
    overlap_slices: int,
) -> dict[str, Any]:
    eff_cfg = cfg_path if cfg_path.is_file() else DEFAULT_STRATEGY_CONFIG_PATH
    _cfg = load_strategy_config_file(eff_cfg)
    apply_strategy_params_to_vrp_backtester_module(_cfg.strategy_params)
    bt_kw: dict[str, Any] = {
        "initial_capital": capital,
        "spy_df": spy_wide,
        "vol_risk_scaling": sc.vol_scaling,
        "r2_crossover_filters": sc.r2_crossover,
        "dd_risk_scaling": sc.dd_risk_scaling,
        "dd_scale_enter": float(dd_enter),
        "dd_scale_exit": float(dd_exit),
        "dd_scale_mult": float(dd_mult),
        "sleeve_risk_fractions": cast(dict, dict(_cfg.sleeve_risk_fractions)),
        "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,
        "overlap_portfolio": sc.overlap_portfolio,
        "overlap_slice_contracts": int(overlap_slices),
        "macro_overlay_enabled": sc.macro_overlay,
        "macro_overlay_total_frac": float(macro_frac),
        "macro_overlay_requires_spy200": True,
        "macro_overlay_exclude_optional_etf": False,
    }

    bt = VRPBacktester(ld, **bt_kw)
    t0 = time.perf_counter()
    bt.run_backtest(trading_days=days, show_progress=False)
    elapsed = time.perf_counter() - t0
    m = bt.metrics()
    return {
        "scenario": asdict(sc),
        "elapsed_sec": round(elapsed, 3),
        "trading_days": len(days),
        "first_day": pd.Timestamp(days[0]).date().isoformat() if days else "",
        "last_day": pd.Timestamp(days[-1]).date().isoformat() if days else "",
        "ending_capital": float(m["ending_capital"]),
        "total_return": float(m["total_return"]),
        "max_drawdown": float(m["max_drawdown"]),
        "cagr": float(m["cagr"]) if isinstance(m["cagr"], float) and math.isfinite(m["cagr"]) else None,
        "total_trades": int(m["total_trades"]),
        "macro_overlay_pnl_usd": float(m.get("macro_overlay_pnl_usd", 0.0)),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="VRP toggle matrix → JSON + Markdown")
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA_DIR)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--strategy-config", type=Path, default=DEFAULT_STRATEGY_CONFIG_PATH)
    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("--max-days", type=int, default=0, help="If >0, only first N days after filters")
    ap.add_argument("--full-macro-factorial", action="store_true", help="32 rows (macro on/off × 16)")
    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("--macro-frac", type=float, default=0.01)
    ap.add_argument("--overlap-slice-contracts", type=int, default=1)
    ap.add_argument("--no-progress", action="store_true", help="Unused (always quiet); for CLI parity")
    ap.add_argument("--json-out", type=Path, default=_JSON_OUT)
    ap.add_argument("--md-out", type=Path, default=_MD_OUT)
    args = ap.parse_args()

    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)

    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)

    cfg_path = args.strategy_config.expanduser()
    scenarios = _scenario_grid(include_macro_factorial=bool(args.full_macro_factorial))
    rows: list[dict[str, Any]] = []
    wall0 = time.perf_counter()
    for i, sc in enumerate(scenarios, 1):
        print(f"[{i}/{len(scenarios)}] {sc.name}", flush=True)
        rows.append(
            _run_one(
                ld,
                spy_wide,
                days,
                float(args.capital),
                cfg_path,
                sc,
                dd_enter=float(args.dd_enter),
                dd_exit=float(args.dd_exit),
                dd_mult=float(args.dd_mult),
                macro_frac=float(args.macro_frac),
                overlap_slices=int(args.overlap_slice_contracts),
            )
        )
    wall = time.perf_counter() - wall0

    payload: dict[str, Any] = {
        "created_at": pd.Timestamp.utcnow().isoformat(),
        "theta_dir": str(theta_dir.resolve()),
        "strategy_config": str(cfg_path.resolve()) if cfg_path.is_file() else str(cfg_path),
        "capital_usd": float(args.capital),
        "filters": {
            "start": args.start or None,
            "end": args.end or None,
            "max_days": int(args.max_days) or None,
        },
        "dd_risk_scaling_params": {
            "enter": float(args.dd_enter),
            "exit": float(args.dd_exit),
            "mult": float(args.dd_mult),
        },
        "macro_overlay_frac_when_enabled": float(args.macro_frac),
        "overlap_slice_contracts": int(args.overlap_slice_contracts),
        "wall_clock_sec": round(wall, 3),
        "rows": rows,
    }

    args.json_out.parent.mkdir(parents=True, exist_ok=True)
    args.json_out.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")

    # Markdown table (sorted by max DD ascending then return descending)
    def _sort_key(r: dict[str, Any]) -> tuple[float, float]:
        return (float(r["max_drawdown"]), -float(r["total_return"]))

    sorted_rows = sorted(rows, key=_sort_key)
    lines: list[str] = [
        "# VRP toggle matrix (machine-generated)",
        "",
        f"- **Artifact:** `{args.json_out.resolve()}` (source of truth)",
        f"- **Generated (UTC):** {payload['created_at']}",
        f"- **Theta dir:** `{payload['theta_dir']}`",
        f"- **Trading days in each run:** {len(days)} ({payload['rows'][0]['first_day']} → {payload['rows'][0]['last_day']})",
        f"- **Capital:** ${args.capital:,.0f}",
        "",
        "## Toggle legend",
        "",
        "| Token | Meaning |",
        "|-------|---------|",
        "| **vol** | `vol_risk_scaling` (VVIX / VIX-momentum size cuts) |",
        "| **r2x** | `r2_crossover_filters` (R2 band: SPY>SMA50, VIX MA/max gate) |",
        "| **dd** | `dd_risk_scaling` (shrink new risk after portfolio DD ≥ {:.0%}; exit ≤ {:.0%}; mult ×{:.2f}) |".format(
            float(args.dd_enter), float(args.dd_exit), float(args.dd_mult)
        ),
        "| **ov** | `overlap_portfolio` (DCA slices; cap still sleeve `target_risk`) |",
        "| **mac** | `macro_overlay` (Tier A daily return bump; frac × capital when ON) |",
        "",
        "**Note:** Metrics are from `VRPBacktester.metrics()` (max DD on `_equity_curve` step points; see engine doc). "
        "Shorter `--max-days` windows are for faster sweeps — **not** comparable to a full-history headline run unless "
        "`--max-days 0`.",
        "",
        "## Results (sorted by max drawdown ↑, then total return ↓)",
        "",
        "| Scenario | End $ | Tot ret | Max DD | CAGR | Trades | Macro PnL $ |",
        "|----------|------:|--------:|-------:|-----:|-------:|------------:|",
    ]
    for r in sorted_rows:
        sc = r["scenario"]
        name = sc["name"].replace("|", "\\|")
        cagr = r["cagr"]
        cagr_s = f"{cagr:.2%}" if isinstance(cagr, float) and math.isfinite(cagr) else "n/a"
        lines.append(
            f"| {name} | {r['ending_capital']:,.2f} | {r['total_return']:.2%} | {r['max_drawdown']:.2%} | "
            f"{cagr_s} | {r['total_trades']} | {r['macro_overlay_pnl_usd']:,.2f} |"
        )
    lines.append("")
    lines.append(f"_Wall clock: {wall:.1f}s for {len(scenarios)} scenarios._")
    lines.append("")
    args.md_out.write_text("\n".join(lines), encoding="utf-8")

    print(f"Wrote {args.json_out.resolve()}", flush=True)
    print(f"Wrote {args.md_out.resolve()}", flush=True)


if __name__ == "__main__":
    main()
