#!/usr/bin/env python3
"""
Steps 1–2 (+ VVIX exit buckets): extended regime metrics, R1 structure digest, portfolio metrics.

Also prints optional VVIX-at-exit quartile PnL (step 7 conditioning probe).

Usage::

    python RenTech/strategy_stack/vrp_research_metrics.py
    python RenTech/strategy_stack/vrp_research_metrics.py --start 2020-01-01 --end 2023-12-31
    python RenTech/strategy_stack/vrp_research_metrics.py --json-out RenTech/data/logs/vrp_metrics.json
"""

from __future__ import annotations

import argparse
import json
import math
import statistics
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 attach_vvix, load_theta_run
from RenTech.strategy_stack.vrp_backtester import (
    DEFAULT_STARTING_CAPITAL,
    VRPBacktester,
    extended_regime_metrics,
    r1_structure_digest,
)


def _vvix_quartile_pnl(trades: list, spy: pd.DataFrame) -> dict:
    if "vvix_close" not in spy.columns:
        return {"skipped": "no vvix_close column"}
    pairs: list[tuple[float, float]] = []
    for t in trades:
        d = pd.Timestamp(t.exit_date).normalize()
        if d not in spy.index:
            continue
        vv = float(spy.loc[d, "vvix_close"])
        if math.isfinite(vv):
            pairs.append((vv, float(t.pnl_usd)))
    if len(pairs) < 8:
        return {"n": len(pairs), "note": "too_few_trades_with_vvix"}
    vals = sorted(x[0] for x in pairs)
    n = len(vals)
    q = [vals[int(n * p)] for p in (0.25, 0.5, 0.75)]
    q1, q2, q3 = q[0], q[1], q[2]
    buckets = {"vvix_q1_lowest": [], "vvix_q2": [], "vvix_q3": [], "vvix_q4_highest": []}
    for vv, pnl in pairs:
        if vv <= q1:
            buckets["vvix_q1_lowest"].append(pnl)
        elif vv <= q2:
            buckets["vvix_q2"].append(pnl)
        elif vv <= q3:
            buckets["vvix_q3"].append(pnl)
        else:
            buckets["vvix_q4_highest"].append(pnl)
    out: dict = {}
    for name, pnls in buckets.items():
        out[name] = {
            "n": len(pnls),
            "mean_pnl": float(statistics.mean(pnls)) if pnls else 0.0,
            "total_pnl": float(sum(pnls)),
        }
    out["vvix_quartile_thresholds"] = {"q1": q1, "q2": q2, "q3": q3}
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description="VRP research: extended metrics + R1 digest (+ VVIX buckets)")
    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("--slippage", type=float, default=None, help="Override slippage factor (default: engine default)")
    ap.add_argument("--no-progress", action="store_true")
    ap.add_argument("--json-out", type=Path, default=None, help="Write combined JSON payload")
    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 in window.")
        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_v = attach_vvix(spy, yf_start, yf_end)

    kw = {"initial_capital": float(args.capital), "spy_df": spy}
    if args.slippage is not None:
        kw["slippage_factor"] = float(args.slippage)
    bt = VRPBacktester(ld, **kw)
    bt.run_backtest(trading_days=days, show_progress=not args.no_progress)

    m = bt.metrics()
    ext = extended_regime_metrics(bt.trade_log, float(args.capital))
    r1 = r1_structure_digest()
    vv = _vvix_quartile_pnl(bt.trade_log, spy_v)

    print("=" * 72)
    print(" VRP research metrics (Theta chunks)")
    print(f" Days: {len(days)}  |  window {d0.date()} → {d1.date()}")
    print("=" * 72)
    print("\n--- R1 structure digest (sanity vs design doc) ---")
    print(json.dumps(r1, indent=2))
    print("\n--- Extended regime metrics ---")
    print(json.dumps(ext, indent=2, default=str))
    print("\n--- Portfolio metrics() ---")
    print(json.dumps({k: m[k] for k in sorted(m.keys())}, indent=2, default=str))
    print("\n--- VVIX at exit (quartile PnL) ---")
    print(json.dumps(vv, indent=2, default=str))

    if args.json_out:
        args.json_out.parent.mkdir(parents=True, exist_ok=True)
        payload = {
            "r1_structure_digest": r1,
            "extended_regime_metrics": ext,
            "metrics": {k: m[k] for k in m},
            "vvix_exit_quartiles": vv,
            "n_days": len(days),
            "d0": str(d0.date()),
            "d1": str(d1.date()),
        }
        args.json_out.write_text(json.dumps(payload, indent=2, default=str), encoding="utf-8")
        print(f"\nWrote {args.json_out}")


if __name__ == "__main__":
    main()
