#!/usr/bin/env python3
"""
Size **VRP main** + optional **IV overlay** sleeves together and report **combined equity**
and **max drawdown**, with optional **daily CSV** for plotting.

Uses trade JSONLs (P&amp;L realized on **exit** date):

* **VRP:** ``vrp_backtest_theta.py --export-trades-jsonl`` → ``pnl_usd``
* **OTM put / straddle:** ``backtest_iv_stress_long_vol.py --out-trades`` → ``pnl_total``
* **Calendar / risk reversal:** ``backtest_iv_calendar_risk_reversal.py --out-trades`` → ``pnl_total``

**Scaling**

* Main daily P&amp;L is multiplied by ``capital_main / vrp_reference_capital`` (default reference
  **100000** to match a typical VRP run).
* Each overlay sleeve daily P&amp;L is multiplied by ``capital_<sleeve> / overlay_reference_capital``
  (default **10000** = “one unit” of the overlay backtest notional).

**Combined equity** = sum of all provided sleeves (each optional).

Example::

    python RenTech/strategy_stack/portfolio_vrp_iv_sleeves.py \\
      --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
      --put-trades RenTech/data/logs/stress_longvol_otm_put.jsonl \\
      --straddle-trades RenTech/data/logs/stress_longvol_straddle.jsonl \\
      --calendar-trades RenTech/data/logs/calendar_spread.jsonl \\
      --risk-reversal-trades RenTech/data/logs/risk_reversal.jsonl \\
      --capital-main 100000 --capital-put 10000 --capital-straddle 10000 \\
      --capital-calendar 10000 --capital-risk-reversal 10000 \\
      --out-equity-csv RenTech/data/logs/portfolio_equity_daily.csv
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from pathlib import Path

import numpy as np
import pandas as pd

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

from RenTech.strategy_stack.iv_mispricing_complement import (
    _load_jsonl,
    _pnl_series_from_trades,
    _require_jsonl,
)


def _max_dd(eq: pd.Series) -> float:
    if eq.empty or len(eq) < 2:
        return 0.0
    peak = eq.cummax()
    x = (peak - eq) / peak.replace(0.0, np.nan)
    return float(x.fillna(0.0).max())


def run_portfolio(
    *,
    vrp_path: Path,
    put_path: Path | None,
    straddle_path: Path | None,
    calendar_path: Path | None,
    risk_reversal_path: Path | None,
    capital_main: float,
    capital_put: float,
    capital_straddle: float,
    capital_calendar: float,
    capital_risk_reversal: float,
    vrp_reference_capital: float,
    overlay_reference_capital: float,
) -> tuple[dict, pd.DataFrame]:
    vrp = _load_jsonl(_require_jsonl(vrp_path, hint="VRP export"))
    put_trades = _load_jsonl(_require_jsonl(put_path, hint="OTM put JSONL")) if put_path else []
    st_trades = _load_jsonl(_require_jsonl(straddle_path, hint="Straddle JSONL")) if straddle_path else []
    cal_trades = _load_jsonl(_require_jsonl(calendar_path, hint="Calendar JSONL")) if calendar_path else []
    rr_trades = _load_jsonl(_require_jsonl(risk_reversal_path, hint="Risk reversal JSONL")) if risk_reversal_path else []

    main = _pnl_series_from_trades(vrp, exit_key="exit_date", pnl_key="pnl_usd")
    put = _pnl_series_from_trades(put_trades, exit_key="exit_date", pnl_key="pnl_total") if put_trades else pd.Series(dtype=float)
    st = _pnl_series_from_trades(st_trades, exit_key="exit_date", pnl_key="pnl_total") if st_trades else pd.Series(dtype=float)
    cal = _pnl_series_from_trades(cal_trades, exit_key="exit_date", pnl_key="pnl_total") if cal_trades else pd.Series(dtype=float)
    rr = _pnl_series_from_trades(rr_trades, exit_key="exit_date", pnl_key="pnl_total") if rr_trades else pd.Series(dtype=float)

    all_idx = (
        main.index.union(put.index).union(st.index).union(cal.index).union(rr.index).sort_values()
    )
    if len(all_idx) == 0:
        raise ValueError("No PnL dates in inputs")

    sm = float(capital_main) / max(float(vrp_reference_capital), 1e-9)
    sp = float(capital_put) / max(float(overlay_reference_capital), 1e-9)
    ss = float(capital_straddle) / max(float(overlay_reference_capital), 1e-9)
    sc = float(capital_calendar) / max(float(overlay_reference_capital), 1e-9)
    sr = float(capital_risk_reversal) / max(float(overlay_reference_capital), 1e-9)

    M = main.reindex(all_idx, fill_value=0.0) * sm
    P = put.reindex(all_idx, fill_value=0.0) * sp
    S = st.reindex(all_idx, fill_value=0.0) * ss
    C = cal.reindex(all_idx, fill_value=0.0) * sc
    R = rr.reindex(all_idx, fill_value=0.0) * sr

    eq_m = float(capital_main) + M.cumsum()
    eq_p = float(capital_put) + P.cumsum()
    eq_s = float(capital_straddle) + S.cumsum()
    eq_c = float(capital_calendar) + C.cumsum()
    eq_r = float(capital_risk_reversal) + R.cumsum()
    total = eq_m + eq_p + eq_s + eq_c + eq_r

    frame = pd.DataFrame(
        {
            "equity_main": eq_m,
            "equity_put": eq_p,
            "equity_straddle": eq_s,
            "equity_calendar": eq_c,
            "equity_risk_reversal": eq_r,
            "equity_total": total,
        },
        index=all_idx,
    )

    meta = {
        "capital_main": float(capital_main),
        "capital_put": float(capital_put),
        "capital_straddle": float(capital_straddle),
        "capital_calendar": float(capital_calendar),
        "capital_risk_reversal": float(capital_risk_reversal),
        "vrp_reference_capital": float(vrp_reference_capital),
        "overlay_reference_capital": float(overlay_reference_capital),
        "ending_main": float(eq_m.iloc[-1]),
        "ending_put": float(eq_p.iloc[-1]),
        "ending_straddle": float(eq_s.iloc[-1]),
        "ending_calendar": float(eq_c.iloc[-1]),
        "ending_risk_reversal": float(eq_r.iloc[-1]),
        "ending_total": float(total.iloc[-1]),
        "max_dd_main": _max_dd(eq_m),
        "max_dd_put": _max_dd(eq_p),
        "max_dd_straddle": _max_dd(eq_s),
        "max_dd_calendar": _max_dd(eq_c),
        "max_dd_risk_reversal": _max_dd(eq_r),
        "max_dd_total": _max_dd(total),
        "total_pnl_main_scaled": float(M.sum()),
        "total_pnl_put_scaled": float(P.sum()),
        "total_pnl_straddle_scaled": float(S.sum()),
        "total_pnl_calendar_scaled": float(C.sum()),
        "total_pnl_risk_reversal_scaled": float(R.sum()),
        "n_days": int(len(all_idx)),
    }
    return meta, frame


def main() -> None:
    ap = argparse.ArgumentParser(description="VRP + optional IV overlay sleeves: equity & DD")
    ap.add_argument("--vrp-trades", type=Path, required=True)
    ap.add_argument("--put-trades", type=Path, default=None, help="OTM put JSONL (optional)")
    ap.add_argument("--straddle-trades", type=Path, default=None, help="Straddle JSONL (optional)")
    ap.add_argument(
        "--calendar-trades",
        type=Path,
        default=None,
        help="Calendar spread JSONL from backtest_iv_calendar_risk_reversal --mode calendar (optional)",
    )
    ap.add_argument(
        "--risk-reversal-trades",
        type=Path,
        default=None,
        help="Risk reversal JSONL from backtest_iv_calendar_risk_reversal --mode risk_reversal (optional)",
    )
    ap.add_argument("--capital-main", type=float, default=100_000.0)
    ap.add_argument("--capital-put", type=float, default=10_000.0)
    ap.add_argument("--capital-straddle", type=float, default=10_000.0)
    ap.add_argument("--capital-calendar", type=float, default=10_000.0)
    ap.add_argument("--capital-risk-reversal", type=float, default=10_000.0)
    ap.add_argument(
        "--vrp-reference-capital",
        type=float,
        default=100_000.0,
        help="VRP export assumed run size (scale main PnL by capital_main / this).",
    )
    ap.add_argument(
        "--overlay-reference-capital",
        type=float,
        default=10_000.0,
        help="One unit of long-vol sleeve notional for scaling.",
    )
    ap.add_argument("--out-equity-csv", type=Path, default=None)
    args = ap.parse_args()

    try:
        meta, frame = run_portfolio(
            vrp_path=args.vrp_trades,
            put_path=args.put_trades,
            straddle_path=args.straddle_trades,
            calendar_path=args.calendar_trades,
            risk_reversal_path=args.risk_reversal_trades,
            capital_main=float(args.capital_main),
            capital_put=float(args.capital_put),
            capital_straddle=float(args.capital_straddle),
            capital_calendar=float(args.capital_calendar),
            capital_risk_reversal=float(args.capital_risk_reversal),
            vrp_reference_capital=float(args.vrp_reference_capital),
            overlay_reference_capital=float(args.overlay_reference_capital),
        )
    except ValueError as e:
        print(f"ERROR: {e}", file=sys.stderr)
        sys.exit(1)

    print(json.dumps(meta, indent=2))

    if args.out_equity_csv:
        p = args.out_equity_csv.expanduser()
        p.parent.mkdir(parents=True, exist_ok=True)
        frame.to_csv(p, date_format="%Y-%m-%d")
        print(f"Wrote daily equity: {p}")


if __name__ == "__main__":
    main()
