#!/usr/bin/env python3
"""
Steps 3–4: slippage stress grid and walk-forward (by calendar year) stability tables.

Usage::

    python RenTech/strategy_stack/vrp_research_sweeps.py --start 2018-01-01 --end 2023-12-31
    python RenTech/strategy_stack/vrp_research_sweeps.py --slippage-only
    python RenTech/strategy_stack/vrp_research_sweeps.py --walkforward-only
"""

from __future__ import annotations

import argparse
import json
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.strategy_stack.vrp_research_common import load_theta_run, year_buckets
from RenTech.strategy_stack.vrp_backtester import (
    DEFAULT_STARTING_CAPITAL,
    VRPBacktester,
    extended_regime_metrics,
)


def _row(m: dict, label: str) -> dict:
    return {
        "label": label,
        "trades": int(m["total_trades"]),
        "total_return": float(m["total_return"]),
        "max_drawdown": float(m["max_drawdown"]),
        "ending_capital": float(m["ending_capital"]),
        "pmcc_trades": int(m["pmcc_trades"]),
        "win_rate": float(m["win_rate"]),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="VRP slippage sweep + walk-forward by year")
    ap.add_argument("--theta-dir", type=Path, default=None)
    ap.add_argument("--capital", type=float, default=DEFAULT_STARTING_CAPITAL)
    ap.add_argument("--start", type=str, default="")
    ap.add_argument("--end", type=str, default="")
    ap.add_argument("--max-days", type=int, default=0)
    ap.add_argument("--no-progress", action="store_true")
    ap.add_argument("--slippage-grid", type=str, default="0.2,0.5,1.0", help="Comma-separated factors")
    ap.add_argument("--slippage-only", action="store_true")
    ap.add_argument("--walkforward-only", action="store_true")
    ap.add_argument("--json-out", type=Path, default=None)
    args = ap.parse_args()

    ld, spy, days, d0, d1 = load_theta_run(
        args.theta_dir,
        start=args.start.strip() or None,
        end=args.end.strip() or None,
    )
    if int(args.max_days) > 0:
        days = days[: int(args.max_days)]
    if not days:
        print("ERROR: no trading days.")
        sys.exit(1)

    cap = float(args.capital)
    show_p = not args.no_progress
    out: dict = {"window": {"d0": str(d0.date()), "d1": str(d1.date()), "n_days": len(days)}}

    do_slip = not args.walkforward_only
    do_wf = not args.slippage_only

    if do_slip:
        grid = [float(x.strip()) for x in args.slippage_grid.split(",") if x.strip()]
        rows = []
        for sf in grid:
            bt = VRPBacktester(ld, initial_capital=cap, spy_df=spy, slippage_factor=sf)
            bt.run_backtest(trading_days=days, show_progress=show_p)
            m = bt.metrics()
            rows.append({"slippage_factor": sf, **_row(m, f"slip={sf}")})
        out["slippage_sweep"] = rows
        print("=" * 72)
        print(" Step 3 — Slippage stress (same days, varying execution friction)")
        print(json.dumps(rows, indent=2))

    if do_wf:
        byy = year_buckets(days)
        wf_rows = []
        for year in sorted(byy.keys()):
            ydays = byy[year]
            if not ydays:
                continue
            bt = VRPBacktester(ld, initial_capital=cap, spy_df=spy)
            bt.run_backtest(trading_days=ydays, show_progress=show_p)
            m = bt.metrics()
            ext = extended_regime_metrics(bt.trade_log, cap)
            wf_rows.append(
                {
                    "year": year,
                    "n_days": len(ydays),
                    **_row(m, str(year)),
                    "regime_n": {k: ext[k]["n"] for k in ext},
                }
            )
        out["walkforward_by_year"] = wf_rows
        print("=" * 72)
        print(" Step 4 — Walk-forward (one backtest per calendar year)")
        print(json.dumps(wf_rows, indent=2))

    if args.json_out:
        args.json_out.parent.mkdir(parents=True, exist_ok=True)
        args.json_out.write_text(json.dumps(out, indent=2), encoding="utf-8")
        print(f"\nWrote {args.json_out}")


if __name__ == "__main__":
    main()
