#!/usr/bin/env python3
"""
**Complement sleeve** for the 4-regime VRP book: align **IV mispricing** trades (e.g. bear call
credit spread JSONL from ``backtest_iv_rich_vol_bear_call_spread.py``) with **VRP** trades exported
via ``vrp_backtest_theta.py --export-trades-jsonl``.

**analyze** — Same calendar: correlation of realized P&amp;L (assigned to **exit** dates), mean
overlay P&amp;L on days the main sleeve lost money, etc.

**combine** — Simulate a two-book portfolio: fixed notionals for main + overlay, optional
**drawdown boost**: overlay realized P&amp;L is scaled by ``1 + dd_boost × drawdown_fraction`` of
the **main** equity path (so the complement sleeve gets larger when VRP is underwater — without
filtering away main trades).

This does **not** guarantee negative correlation; use **analyze** first. Typical use: run VRP with
``--export-trades-jsonl``, run IV overlay backtest with ``--out-trades``, then this script.

Example::

    python RenTech/strategy_stack/vrp_backtest_theta.py --export-trades-jsonl RenTech/data/logs/vrp_trades.jsonl
    python RenTech/strategy_stack/backtest_iv_rich_vol_bear_call_spread.py ... --out-trades RenTech/data/logs/bear_spread.jsonl
    # Or use straddle JSONL: --overlay-trades RenTech/data/logs/iv_mispricing_straddle_trades.jsonl
    python RenTech/strategy_stack/iv_mispricing_complement.py analyze \\
        --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
        --overlay-trades RenTech/data/logs/bear_spread.jsonl
    python RenTech/strategy_stack/iv_mispricing_complement.py combine \\
        --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
        --overlay-trades RenTech/data/logs/bear_spread.jsonl \\
        --capital-main 100000 --capital-overlay 25000 --dd-boost 1.5
"""

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))


def _require_jsonl(path: Path, *, hint: str) -> Path:
    p = path.expanduser().resolve()
    if not p.is_file():
        print(
            f"ERROR: file not found: {p}\n"
            f"  ({hint})",
            file=sys.stderr,
        )
        sys.exit(1)
    return p


def _load_jsonl(path: Path) -> list[dict]:
    rows: list[dict] = []
    with path.open(encoding="utf-8") as f:
        for line in f:
            line = line.strip()
            if not line:
                continue
            rows.append(json.loads(line))
    return rows


def _pnl_series_from_trades(
    trades: list[dict],
    *,
    exit_key: str,
    pnl_key: str,
) -> pd.Series:
    """Daily P&amp;L summed on **exit** date (normalized midnight)."""
    s: dict[pd.Timestamp, float] = {}
    for r in trades:
        ex = r.get(exit_key)
        if ex is None:
            continue
        d = pd.Timestamp(ex).normalize()
        p = float(r.get(pnl_key, 0.0) or 0.0)
        if not math.isfinite(p):
            continue
        s[d] = s.get(d, 0.0) + p
    if not s:
        return pd.Series(dtype=float)
    idx = pd.DatetimeIndex(sorted(s.keys()))
    return pd.Series([s[k] for k in idx], index=idx, dtype=float)


def run_analyze(vrp_path: Path, overlay_path: Path) -> dict:
    vrp = _load_jsonl(_require_jsonl(vrp_path, hint="Export VRP with: vrp_backtest_theta.py --export-trades-jsonl ..."))
    ov = _load_jsonl(
        _require_jsonl(
            overlay_path,
            hint="Create overlay trades, e.g. backtest_iv_rich_vol_bear_call_spread.py ... --out-trades RenTech/data/logs/bear_spread.jsonl "
            "or use iv_mispricing_straddle_trades.jsonl from the straddle backtest.",
        )
    )
    main = _pnl_series_from_trades(vrp, exit_key="exit_date", pnl_key="pnl_usd")
    ovl = _pnl_series_from_trades(ov, exit_key="exit_date", pnl_key="pnl_total")

    all_idx = main.index.union(ovl.index).sort_values()
    M = main.reindex(all_idx, fill_value=0.0)
    O = ovl.reindex(all_idx, fill_value=0.0)

    mask_both = (M != 0) | (O != 0)
    if int(mask_both.sum()) < 5:
        return {"error": "too few overlapping days with PnL"}

    Mc = M[mask_both]
    Oc = O[mask_both]
    corr = float(Mc.corr(Oc)) if Mc.std() > 0 and Oc.std() > 0 else float("nan")

    main_loss = M < 0
    n_loss = int(main_loss.sum())
    mean_ov_when_main_loss = float(O[main_loss].mean()) if n_loss else float("nan")
    mean_ov_when_main_win = float(O[~main_loss].mean()) if int((~main_loss).sum()) else float("nan")

    return {
        "n_vrp_trades": len(vrp),
        "n_overlay_trades": len(ov),
        "calendar_days_with_any_pnl": int(mask_both.sum()),
        "corr_daily_pnl_exit_day": corr,
        "days_main_pnl_negative": n_loss,
        "mean_overlay_pnl_when_main_negative": mean_ov_when_main_loss,
        "mean_overlay_pnl_when_main_non_negative": mean_ov_when_main_win,
        "total_main_pnl": float(M.sum()),
        "total_overlay_pnl": float(O.sum()),
    }


def run_combine(
    vrp_path: Path,
    overlay_path: Path,
    *,
    capital_main: float,
    capital_overlay: float,
    dd_boost: float,
) -> dict:
    """
    Cumulative wealth = main book + overlay book. Overlay daily P&amp;L scaled by
    ``1 + dd_boost * dd_main`` where ``dd_main`` is drawdown of main equity (prev day).
    """
    vrp = _load_jsonl(_require_jsonl(vrp_path, hint="Export VRP with --export-trades-jsonl"))
    ov = _load_jsonl(
        _require_jsonl(
            overlay_path,
            hint="Create overlay JSONL (bear spread or straddle backtest --out-trades).",
        )
    )
    main = _pnl_series_from_trades(vrp, exit_key="exit_date", pnl_key="pnl_usd")
    ovl = _pnl_series_from_trades(ov, exit_key="exit_date", pnl_key="pnl_total")

    all_idx = main.index.union(ovl.index).sort_values()
    M = main.reindex(all_idx, fill_value=0.0)
    O_raw = ovl.reindex(all_idx, fill_value=0.0)

    eq_m = float(capital_main) + M.cumsum()
    peak = eq_m.cummax()
    dd = (peak - eq_m) / peak.replace(0.0, np.nan)
    dd = dd.fillna(0.0).clip(0.0, 1.0)
    scale = 1.0 + float(dd_boost) * dd.shift(1).fillna(0.0)
    O = O_raw * scale
    eq_o = float(capital_overlay) + O.cumsum()
    total = eq_m + eq_o

    def _max_dd(eq: pd.Series) -> float:
        p = eq.cummax()
        x = (p - eq) / p.replace(0.0, np.nan)
        return float(x.max()) if len(eq) else 0.0

    return {
        "capital_main": float(capital_main),
        "capital_overlay": float(capital_overlay),
        "dd_boost": float(dd_boost),
        "ending_main": float(eq_m.iloc[-1]) if len(eq_m) else float(capital_main),
        "ending_overlay": float(eq_o.iloc[-1]) if len(eq_o) else float(capital_overlay),
        "ending_combined": float(total.iloc[-1]) if len(total) else float(capital_main + capital_overlay),
        "max_dd_main": _max_dd(eq_m),
        "max_dd_overlay": _max_dd(eq_o),
        "max_dd_combined": _max_dd(total),
        "total_pnl_main": float(M.sum()),
        "total_pnl_overlay_scaled": float(O.sum()),
    }


def main() -> None:
    ap = argparse.ArgumentParser(description="VRP + IV mispricing complement sleeve analysis")
    sub = ap.add_subparsers(dest="cmd", required=True)

    a = sub.add_parser("analyze", help="Correlation and conditional means")
    a.add_argument("--vrp-trades", type=Path, required=True)
    a.add_argument("--overlay-trades", type=Path, required=True)

    c = sub.add_parser("combine", help="Simulate combined equity with optional DD boost on overlay")
    c.add_argument("--vrp-trades", type=Path, required=True)
    c.add_argument("--overlay-trades", type=Path, required=True)
    c.add_argument("--capital-main", type=float, default=100_000.0)
    c.add_argument("--capital-overlay", type=float, default=25_000.0)
    c.add_argument(
        "--dd-boost",
        type=float,
        default=0.0,
        help="Multiply overlay PnL by (1 + dd_boost * main_drawdown); 0 = fixed overlay size.",
    )

    args = ap.parse_args()
    if args.cmd == "analyze":
        out = run_analyze(args.vrp_trades, args.overlay_trades)
        print(json.dumps(out, indent=2))
        if "error" in out:
            sys.exit(1)
    elif args.cmd == "combine":
        out = run_combine(
            args.vrp_trades,
            args.overlay_trades,
            capital_main=float(args.capital_main),
            capital_overlay=float(args.capital_overlay),
            dd_boost=float(args.dd_boost),
        )
        print(json.dumps(out, indent=2))


if __name__ == "__main__":
    main()
