#!/usr/bin/env python3
"""
Merge **VRP** (options JSONL) + **IV engine** (put/straddle/RR) + **lit4** (legacy put + RR JSONLs)
+ **VXX** (bear + long) on **one** calendar, with explicit per-sleeve capital/ref scaling.

``portfolio_vrp_plus_vxx`` only has **one** put slot and **one** RR slot, so **IV and lit4 cannot**
both occupy put/RR there. This script adds lit4 as **extra** daily PnL on top of the IV+VXX merge
(same scale law as the merge: ``scaled_pnl = raw * (capital / ref)`` per sleeve).

Writes per-sleeve scaled CSVs, ``{prefix}_ALL_SLEEVES_trade_log.csv`` (sorted by exit date),
``{prefix}_equity_daily.csv``, and ``{prefix}_manifest.txt``.

Default sleeve budgets match the interior Sharpe-opt IV book (same as ``export_sharpe_opt_fullbook_trades.py``).
lit4 defaults: **$1** capital vs **$1** ref each sleeve (raw 1-lot cumulative PnL on the book).
"""
from __future__ import annotations

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

_REPO_ROOT = Path(__file__).resolve().parents[2]
_LOGS = _REPO_ROOT / "RenTech" / "data" / "logs"

if str(_REPO_ROOT) not in sys.path:
    sys.path.insert(0, str(_REPO_ROOT))

import pandas as pd

from RenTech.strategy_stack.iv_mispricing_complement import _load_jsonl, _pnl_series_from_trades
from RenTech.strategy_stack.portfolio_vrp_plus_vxx import (
    _infer_broker_risk_overlay,
    _infer_broker_risk_vxx,
    _metrics_block,
    _overlay_sleeve_risk_reference_usd,
    load_portfolio_pnl_basis,
)

_DEFAULT_PUT_FRAC = 0.00523407
_DEFAULT_STRADDLE_FRAC = 0.05749554
_DEFAULT_RR_FRAC = 0.03388598
_DEFAULT_VXX_FRAC = 0.23201529
_DEFAULT_VRP_JSONL = _LOGS / "vrp_low_dd_ov2_vxxbundle_vrp_trades.jsonl"


def _scale_rows(
    rows: list[dict],
    *,
    sleeve: str,
    pnl_key: str,
    capital_usd: float,
    ref_usd: float,
    exit_key: str = "exit_date",
) -> list[dict]:
    ref = max(float(ref_usd), 1e-12)
    s = float(capital_usd) / ref
    out: list[dict] = []
    for r in rows:
        raw = float(r.get(pnl_key, 0.0) or 0.0)
        ex = r.get(exit_key)
        en = r.get("entry_date", "")
        out.append(
            {
                "sleeve": sleeve,
                "entry_date": str(en)[:10] if en is not None else "",
                "exit_date": str(ex)[:10] if ex is not None else "",
                "pnl_usd_scaled": round(raw * s, 6),
                "pnl_usd_raw_in_log": round(raw, 6),
                "sleeve_scale_factor": round(s, 10),
                "sleeve_risk_ref_usd": round(ref, 2),
                "sleeve_capital_usd": round(float(capital_usd), 2),
            }
        )
    return out


def _vrp_rows_from_jsonl(path: Path, *, sleeve: str, capital_usd: float, ref_usd: float) -> list[dict]:
    rows = _load_jsonl(path)
    ref = max(float(ref_usd), 1e-12)
    s = float(capital_usd) / ref
    out: list[dict] = []
    for r in rows:
        raw = float(r.get("pnl_usd", 0.0) or 0.0)
        ex = r.get("exit_date")
        en = r.get("entry_date", "")
        out.append(
            {
                "sleeve": sleeve,
                "entry_date": str(en)[:10] if en is not None else "",
                "exit_date": str(ex)[:10] if ex is not None else "",
                "pnl_usd_scaled": round(raw * s, 6),
                "pnl_usd_raw_in_log": round(raw, 6),
                "sleeve_scale_factor": round(s, 10),
                "sleeve_risk_ref_usd": round(ref, 2),
                "sleeve_capital_usd": round(float(capital_usd), 2),
                "regime": str(r.get("regime", "")),
                "qty": int(r.get("qty", 0) or 0),
                "exit_reason": str(r.get("exit_reason", "")),
                "legs_json": str(r.get("legs_json", "") or ""),
            }
        )
    return out


def _write_simple_csv(path: Path, fieldnames: list[str], rows: list[dict]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    with path.open("w", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=fieldnames, extrasaction="ignore")
        w.writeheader()
        for r in rows:
            w.writerow(r)


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--vrp-ref", type=float, default=100_000.0)
    ap.add_argument("--capital-vrp-pct", type=float, default=1.0, metavar="FRAC")
    ap.add_argument("--capital-put-pct", type=float, default=_DEFAULT_PUT_FRAC)
    ap.add_argument("--capital-straddle-pct", type=float, default=_DEFAULT_STRADDLE_FRAC)
    ap.add_argument("--capital-risk-reversal-pct", type=float, default=_DEFAULT_RR_FRAC)
    ap.add_argument("--capital-vxx-pct", type=float, default=_DEFAULT_VXX_FRAC)
    ap.add_argument("--vxx-bear-pct", type=float, default=90.0)
    ap.add_argument("--vxx-call-pct", type=float, default=10.0)
    ap.add_argument("--vrp-trades", type=Path, default=_DEFAULT_VRP_JSONL)
    ap.add_argument("--put-trades", type=Path, default=_LOGS / "engine_iv_otm_put_ivx1.jsonl")
    ap.add_argument("--straddle-trades", type=Path, default=_LOGS / "engine_iv_straddle_ivx1.jsonl")
    ap.add_argument("--risk-reversal-trades", type=Path, default=_LOGS / "engine_iv_risk_reversal_ivx1.jsonl")
    ap.add_argument(
        "--vxx-bear-trades",
        type=Path,
        default=_LOGS / "overall_portfolio_trade_audits/engine_vxx_1pct_2016_2026/vxx_portfolio_bear_call.jsonl",
    )
    ap.add_argument(
        "--vxx-call-trades",
        type=Path,
        default=_LOGS / "overall_portfolio_trade_audits/engine_vxx_1pct_2016_2026/vxx_portfolio_long_call.jsonl",
    )
    ap.add_argument("--lit4-put-trades", type=Path, default=_LOGS / "legacy4_put_overlay.jsonl")
    ap.add_argument("--lit4-rr-trades", type=Path, default=_LOGS / "legacy4_rr_overlay.jsonl")
    ap.add_argument("--lit4-put-capital", type=float, default=1.0, help="USD risk budget for lit4 put sleeve")
    ap.add_argument("--lit4-put-ref", type=float, default=1.0, help="Ref USD for lit4 put (default 1 = raw 1-lot)")
    ap.add_argument("--lit4-rr-capital", type=float, default=1.0)
    ap.add_argument("--lit4-rr-ref", type=float, default=1.0)
    ap.add_argument("--out-dir", type=Path, default=_LOGS)
    ap.add_argument("--out-prefix", type=str, default="vrp_iv_lit4_vxx_10dd_merged")
    args = ap.parse_args()

    cap = float(args.capital)
    prefix = str(args.out_prefix).strip() or "vrp_iv_lit4_vxx"
    out_dir = args.out_dir.expanduser().resolve()
    out_dir.mkdir(parents=True, exist_ok=True)

    cap_vrp = cap * float(args.capital_vrp_pct)
    cap_put = cap * float(args.capital_put_pct)
    cap_str = cap * float(args.capital_straddle_pct)
    cap_rr = cap * float(args.capital_risk_reversal_pct)
    cap_vxx = cap * float(args.capital_vxx_pct)
    bear_usd = cap_vxx * (float(args.vxx_bear_pct) / 100.0)
    call_usd = cap_vxx * (float(args.vxx_call_pct) / 100.0)

    vrp_path = args.vrp_trades.expanduser().resolve()
    put_p = args.put_trades.expanduser().resolve()
    str_p = args.straddle_trades.expanduser().resolve()
    rr_p = args.risk_reversal_trades.expanduser().resolve()
    vb_p = args.vxx_bear_trades.expanduser().resolve()
    vc_p = args.vxx_call_trades.expanduser().resolve()
    lit4_p = args.lit4_put_trades.expanduser().resolve()
    lit4_r = args.lit4_rr_trades.expanduser().resolve()

    ref_put = _overlay_sleeve_risk_reference_usd(put_p, _infer_broker_risk_overlay)
    ref_str = _overlay_sleeve_risk_reference_usd(str_p, _infer_broker_risk_overlay)
    ref_rr = _overlay_sleeve_risk_reference_usd(rr_p, _infer_broker_risk_overlay)
    ref_vb = _overlay_sleeve_risk_reference_usd(vb_p, _infer_broker_risk_vxx)
    ref_vc = _overlay_sleeve_risk_reference_usd(vc_p, _infer_broker_risk_vxx)
    ref_l4p = max(float(args.lit4_put_ref), 1e-12)
    ref_l4r = max(float(args.lit4_rr_ref), 1e-12)

    basis = load_portfolio_pnl_basis(
        vrp_trades=vrp_path,
        vrp_ref=float(args.vrp_ref),
        put_trades=put_p,
        straddle_trades=str_p,
        risk_reversal_trades=rr_p,
        vxx_bear_trades=vb_p,
        vxx_call_trades=vc_p,
        total_portfolio_capital=cap,
    )
    idx = basis.index
    rows_lp = _load_jsonl(lit4_p) if lit4_p.is_file() else []
    rows_lr = _load_jsonl(lit4_r) if lit4_r.is_file() else []
    raw_lp = _pnl_series_from_trades(rows_lp, exit_key="exit_date", pnl_key="pnl_total")
    raw_lr = _pnl_series_from_trades(rows_lr, exit_key="exit_date", pnl_key="pnl_total")
    # Include lit4 exit dates that extend past the VRP/IV/VXX union so RR tails are not dropped.
    idx = idx.union(raw_lp.index).union(raw_lr.index).sort_values()
    B_v = basis.B_v.reindex(idx, fill_value=0.0)
    B_p = basis.B_p.reindex(idx, fill_value=0.0)
    B_s = basis.B_s.reindex(idx, fill_value=0.0)
    B_r = basis.B_r.reindex(idx, fill_value=0.0)
    B_vb = basis.B_vb.reindex(idx, fill_value=0.0)
    B_vc = basis.B_vc.reindex(idx, fill_value=0.0)
    eq0 = basis.eq0
    pr, sr, rr_ref = max(basis.put_risk_ref, 1e-12), max(basis.straddle_risk_ref, 1e-12), max(basis.rr_risk_ref, 1e-12)
    vbr, vcr = max(basis.vxx_bear_risk_ref, 1e-12), max(basis.vxx_call_risk_ref, 1e-12)

    vrp_num = float(cap_vrp)
    V = vrp_num * B_v
    P = (cap_put / pr) * B_p
    S = (cap_str / sr) * B_s
    R = (cap_rr / rr_ref) * B_r
    VB = (bear_usd / vbr) * B_vb
    VC = (call_usd / vcr) * B_vc

    s_lp = raw_lp.reindex(idx, fill_value=0.0)
    s_lr = raw_lr.reindex(idx, fill_value=0.0)
    Lp = (float(args.lit4_put_capital) / ref_l4p) * s_lp
    Lr = (float(args.lit4_rr_capital) / ref_l4r) * s_lr

    pnl_full = V + P + S + R + VB + VC + Lp + Lr
    eq_full = eq0 + pnl_full.cumsum()
    m_full = _metrics_block(eq_full, "VRP+IV+lit4+VXX")

    eq_csv = out_dir / f"{prefix}_equity_daily.csv"
    pd.DataFrame(
        {
            "eq_full_portfolio": eq_full,
            "pnl_vrp": V,
            "pnl_iv_put": P,
            "pnl_iv_straddle": S,
            "pnl_iv_rr": R,
            "pnl_lit4_put": Lp,
            "pnl_lit4_rr": Lr,
            "pnl_vxx_bear": VB,
            "pnl_vxx_call": VC,
            "pnl_total_day": pnl_full,
        },
        index=idx,
    ).to_csv(eq_csv)

    vrp_flat = _vrp_rows_from_jsonl(vrp_path, sleeve="vrp", capital_usd=cap_vrp, ref_usd=float(args.vrp_ref))
    scaled_put = _scale_rows(_load_jsonl(put_p), sleeve="iv_otm_put", pnl_key="pnl_total", capital_usd=cap_put, ref_usd=ref_put)
    scaled_str = _scale_rows(_load_jsonl(str_p), sleeve="iv_straddle", pnl_key="pnl_total", capital_usd=cap_str, ref_usd=ref_str)
    scaled_rr = _scale_rows(_load_jsonl(rr_p), sleeve="iv_risk_reversal", pnl_key="pnl_total", capital_usd=cap_rr, ref_usd=ref_rr)
    scaled_vb = _scale_rows(_load_jsonl(vb_p), sleeve="vxx_bear_call", pnl_key="pnl_total", capital_usd=bear_usd, ref_usd=ref_vb)
    scaled_vc = _scale_rows(_load_jsonl(vc_p), sleeve="vxx_long_call", pnl_key="pnl_total", capital_usd=call_usd, ref_usd=ref_vc)
    scaled_l4p = _scale_rows(
        rows_lp,
        sleeve="lit4_put",
        pnl_key="pnl_total",
        capital_usd=float(args.lit4_put_capital),
        ref_usd=ref_l4p,
    )
    scaled_l4r = _scale_rows(
        rows_lr,
        sleeve="lit4_rr",
        pnl_key="pnl_total",
        capital_usd=float(args.lit4_rr_capital),
        ref_usd=ref_l4r,
    )

    combined = vrp_flat + scaled_put + scaled_str + scaled_rr + scaled_vb + scaled_vc + scaled_l4p + scaled_l4r

    def _sort_key(r: dict) -> tuple:
        return (r.get("exit_date") or "", str(r.get("sleeve") or ""))

    combined.sort(key=_sort_key)
    comb_fields = sorted({k for r in combined for k in r})
    comb_path = out_dir / f"{prefix}_ALL_SLEEVES_trade_log.csv"
    _write_simple_csv(comb_path, comb_fields, combined)

    base_cols = [
        "sleeve",
        "entry_date",
        "exit_date",
        "pnl_usd_scaled",
        "pnl_usd_raw_in_log",
        "sleeve_scale_factor",
        "sleeve_risk_ref_usd",
        "sleeve_capital_usd",
    ]
    _write_simple_csv(out_dir / f"{prefix}_iv_put_scaled.csv", base_cols, scaled_put)
    _write_simple_csv(out_dir / f"{prefix}_iv_straddle_scaled.csv", base_cols, scaled_str)
    _write_simple_csv(out_dir / f"{prefix}_iv_rr_scaled.csv", base_cols, scaled_rr)
    _write_simple_csv(out_dir / f"{prefix}_vxx_bear_scaled.csv", base_cols, scaled_vb)
    _write_simple_csv(out_dir / f"{prefix}_vxx_long_scaled.csv", base_cols, scaled_vc)
    _write_simple_csv(out_dir / f"{prefix}_lit4_put_scaled.csv", base_cols, scaled_l4p)
    _write_simple_csv(out_dir / f"{prefix}_lit4_rr_scaled.csv", base_cols, scaled_l4r)
    vrp_cols = [c for c in vrp_flat[0].keys()] if vrp_flat else base_cols
    _write_simple_csv(out_dir / f"{prefix}_vrp_scaled.csv", vrp_cols, vrp_flat)

    manifest = out_dir / f"{prefix}_manifest.txt"
    lines = [
        "VRP + IV engine + lit4 + VXX (lit4 added on top of IV merge calendar; see script docstring).",
        f"prefix: {prefix}",
        f"capital: {cap:g}",
        f"vrp_trades: {vrp_path}",
        f"VRP: capital={cap_vrp:,.2f}  ref={float(args.vrp_ref):,.2f}",
        f"IV put: capital={cap_put:,.2f}  ref={ref_put:,.2f}",
        f"IV straddle: capital={cap_str:,.2f}  ref={ref_str:,.2f}",
        f"IV rr: capital={cap_rr:,.2f}  ref={ref_rr:,.2f}",
        f"VXX: total_capital={cap_vxx:,.2f}  bear={bear_usd:,.2f}  call={call_usd:,.2f}",
        f"lit4 put: capital={float(args.lit4_put_capital):,.2f}  ref={ref_l4p:,.2f}  n_trades={len(scaled_l4p)}",
        f"lit4 rr: capital={float(args.lit4_rr_capital):,.2f}  ref={ref_l4r:,.2f}  n_trades={len(scaled_l4r)}",
        "",
        f"calendar: {idx[0].date()} → {idx[-1].date()}  days={len(idx)}",
        f"FULL return%={m_full['return_pct']:.2f}  maxDD%={m_full['max_dd_pct']:.2f}  CAGR%={m_full['cagr_pct']:.2f}  Sharpe={m_full['sharpe']:.3f}  end=${m_full['end_equity']:,.0f}",
        "",
        "outputs:",
        f"  {eq_csv}",
        f"  {comb_path}",
    ]
    manifest.write_text("\n".join(lines) + "\n", encoding="utf-8")

    print("\n".join(lines))
    print(f"\nWrote {comb_path}")


if __name__ == "__main__":
    main()
