#!/usr/bin/env python3
"""
Optimize **per-overlay risk as a fraction of total account capital** for the merged book:

  VRP (fixed scale) + IV overlays (OTM put, straddle, risk reversal) + VXX sleeve.

This wraps the existing optimizers in ``portfolio_vrp_plus_vxx.py``:

* ``--mode sharpe-frac`` → :func:`optimize_portfolio_sharpe` (maximize Sharpe on fractions).
* ``--mode calmar-usd`` → :func:`optimize_sleeve_notionals` (maximize Calmar + w·Sharpe on USD budgets),
  then converts budgets to fractions of ``--total-capital``.

Output JSON is suitable to paste under ``overlay_risk_fractions`` / portfolio CLI flags.

Examples::

    python RenTech/strategy_stack/optimize_overlay_risk_fracs.py \\
      --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
      --total-capital 100000 \\
      --out-json RenTech/strategy_stack/overlay_risk_fracs_optimized.json

    python RenTech/strategy_stack/optimize_overlay_risk_fracs.py \\
      --mode calmar-usd --max-dd-pct 12 --notionals-budget-pct 0.25
"""

from __future__ import annotations

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

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

from RenTech.strategy_stack.portfolio_vrp_plus_vxx import (
    LOGS,
    optimize_portfolio_sharpe,
    optimize_sleeve_notionals,
)


def _write_json(path: Path, payload: dict[str, Any]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, indent=2), encoding="utf-8")


def main() -> None:
    ap = argparse.ArgumentParser(
        description="Optimize overlay+VXX risk fractions vs VRP trade JSONL (portfolio merge model)."
    )
    ap.add_argument(
        "--vrp-trades",
        type=Path,
        default=LOGS / "vrp_trades.jsonl",
        help="VRP trades JSONL (default: RenTech/data/logs/vrp_trades.jsonl).",
    )
    ap.add_argument("--total-capital", type=float, default=100_000.0)
    ap.add_argument(
        "--capital-vrp-pct",
        type=float,
        default=None,
        metavar="FRAC",
        help="If set, VRP scale = total_capital * frac (else full book = total_capital for VRP leg).",
    )
    ap.add_argument(
        "--mode",
        choices=("sharpe-frac", "calmar-usd"),
        default="sharpe-frac",
        help="Optimization objective (default: sharpe-frac).",
    )
    ap.add_argument(
        "--sum-cap-pct",
        type=float,
        default=0.35,
        help="For sharpe-frac: cap on put+straddle+rr+vxx fractions (0 disables). Default 0.35.",
    )
    ap.add_argument(
        "--max-sleeve-pct",
        type=float,
        default=0.25,
        help="For sharpe-frac: max fraction per sleeve. Default 0.25.",
    )
    ap.add_argument(
        "--max-dd-pct",
        type=float,
        default=10.0,
        help="For sharpe-frac: hard max DD %% (0 disables). For calmar-usd: soft DD target. Default 10.",
    )
    ap.add_argument(
        "--notionals-budget-pct",
        type=float,
        default=0.25,
        help="For calmar-usd: combined USD cap on overlay+VXX as fraction of total capital (0 = no cap).",
    )
    ap.add_argument(
        "--max-single-sleeve-pct",
        type=float,
        default=0.15,
        help="For calmar-usd: max USD per sleeve as fraction of total capital (mapped to USD cap).",
    )
    ap.add_argument("--sharpe-weight", type=float, default=0.25, help="For calmar-usd objective only.")
    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-ref", type=float, default=100_000.0)
    ap.add_argument("--opt-maxiter", type=int, default=80)
    ap.add_argument("--opt-seed", type=int, default=0)
    ap.add_argument("--polish", action="store_true")
    ap.add_argument(
        "--out-json",
        type=Path,
        default=_REPO / "RenTech" / "strategy_stack" / "overlay_risk_fracs_optimized.json",
    )
    ap.add_argument("--put-trades", type=Path, default=None)
    ap.add_argument("--straddle-trades", type=Path, default=None)
    ap.add_argument("--risk-reversal-trades", type=Path, default=None)
    ap.add_argument("--vxx-bear-trades", type=Path, default=None)
    ap.add_argument("--vxx-call-trades", type=Path, default=None)
    args = ap.parse_args()

    vrp_path = args.vrp_trades.expanduser().resolve()
    if not vrp_path.is_file():
        print(f"ERROR: VRP trades not found: {vrp_path}", file=sys.stderr)
        sys.exit(1)

    tc = float(args.total_capital)
    cap_vrp = tc * float(args.capital_vrp_pct) if args.capital_vrp_pct is not None else None

    kw_paths: dict[str, Any] = {}
    if args.put_trades is not None:
        kw_paths["put_trades"] = args.put_trades.expanduser().resolve()
    if args.straddle_trades is not None:
        kw_paths["straddle_trades"] = args.straddle_trades.expanduser().resolve()
    if args.risk_reversal_trades is not None:
        kw_paths["risk_reversal_trades"] = args.risk_reversal_trades.expanduser().resolve()
    if args.vxx_bear_trades is not None:
        kw_paths["vxx_bear_trades"] = args.vxx_bear_trades.expanduser().resolve()
    if args.vxx_call_trades is not None:
        kw_paths["vxx_call_trades"] = args.vxx_call_trades.expanduser().resolve()

    if str(args.mode) == "sharpe-frac":
        sum_cap = None if float(args.sum_cap_pct) <= 0 else float(args.sum_cap_pct)
        dd_cap = None if float(args.max_dd_pct) <= 0 else float(args.max_dd_pct)
        r = optimize_portfolio_sharpe(
            vrp_path,
            total_portfolio_capital=tc,
            capital_vrp=cap_vrp,
            sum_risk_budget_pct=sum_cap,
            max_sleeve_pct=float(args.max_sleeve_pct),
            max_dd_limit_pct=dd_cap,
            vxx_bear_pct=float(args.vxx_bear_pct),
            vxx_call_pct=float(args.vxx_call_pct),
            vrp_ref=float(args.vrp_ref),
            maxiter=int(args.opt_maxiter),
            seed=int(args.opt_seed),
            polish=bool(args.polish),
            **kw_paths,
        )
        out: dict[str, Any] = {
            "mode": "sharpe-frac",
            "vrp_trades": str(vrp_path),
            "total_capital": tc,
            "feasible_dd": r["feasible_dd"],
            "max_drawdown_pct": r["max_drawdown_pct"],
            "sharpe": r["sharpe"],
            "calmar": r["calmar"],
            "cagr_pct": r["cagr_pct"],
            "return_pct": r["return_pct"],
            "end_equity": r["end_equity"],
            "overlay_risk_fractions": {
                "stress_longvol_otm_put": r["put_frac"],
                "stress_longvol_straddle": r["straddle_frac"],
                "risk_reversal": r["rr_frac"],
                "vxx_sleeve": r["vxx_frac"],
            },
            "sum_overlay_fractions": r["sum_frac"],
            "portfolio_vrp_plus_vxx_cli_snippet": {
                "total_capital": tc,
                "capital_put_pct": r["put_frac"],
                "capital_straddle_pct": r["straddle_frac"],
                "capital_risk_reversal_pct": r["rr_frac"],
                "capital_vxx_pct": r["vxx_frac"],
                "vxx_bear_pct": float(args.vxx_bear_pct),
                "vxx_call_pct": float(args.vxx_call_pct),
            },
            "scipy_success": r["scipy_success"],
            "scipy_message": r["scipy_message"],
            "nit": r["nit"],
        }
    else:
        budget = None if float(args.notionals_budget_pct) <= 0 else tc * float(args.notionals_budget_pct)
        max_one = tc * float(args.max_single_sleeve_pct)
        r = optimize_sleeve_notionals(
            vrp_path,
            total_portfolio_capital=tc,
            capital_vrp=cap_vrp,
            max_dd_limit_pct=float(args.max_dd_pct),
            additional_notionals_cap=budget,
            max_single_sleeve=max_one,
            sharpe_weight=float(args.sharpe_weight),
            vxx_bear_pct=float(args.vxx_bear_pct),
            vxx_call_pct=float(args.vxx_call_pct),
            vrp_ref=float(args.vrp_ref),
            maxiter=int(args.opt_maxiter),
            seed=int(args.opt_seed),
            polish=bool(args.polish),
            **kw_paths,
        )
        cp, cs, cr, cv = r["capital_put"], r["capital_straddle"], r["capital_risk_reversal"], r["capital_vxx"]
        out = {
            "mode": "calmar-usd",
            "vrp_trades": str(vrp_path),
            "total_capital": tc,
            "feasible": r["feasible"],
            "max_drawdown_pct": r["max_drawdown_pct"],
            "sharpe": r["sharpe"],
            "calmar": r["calmar"],
            "cagr_pct": r["cagr_pct"],
            "return_pct": r["return_pct"],
            "end_equity": r["end_equity"],
            "overlay_risk_fractions": {
                "stress_longvol_otm_put": cp / tc,
                "stress_longvol_straddle": cs / tc,
                "risk_reversal": cr / tc,
                "vxx_sleeve": cv / tc,
            },
            "overlay_risk_usd": {
                "stress_longvol_otm_put": cp,
                "stress_longvol_straddle": cs,
                "risk_reversal": cr,
                "vxx_sleeve": cv,
            },
            "sum_overlay_fractions": (cp + cs + cr + cv) / tc,
            "objective_score": r["objective_score"],
            "scipy_success": r["scipy_success"],
            "scipy_message": r["scipy_message"],
            "nit": r["nit"],
        }

    _write_json(args.out_json.expanduser(), out)

    o = out["overlay_risk_fractions"]
    print("=" * 72)
    print(f"Optimized overlay fractions ({args.mode})  →  {args.out_json}")
    print("=" * 72)
    for k, v in o.items():
        print(f"  {k}: {float(v):.6f}  (${float(v) * tc:,.0f} @ ${tc:,.0f} book)")
    print(f"  sum(overlay fracs): {float(out['sum_overlay_fractions']):.6f}")
    _cg = out.get("cagr_pct")
    _cgs = f"{_cg:.2f}%" if isinstance(_cg, (int, float)) and math.isfinite(float(_cg)) else "n/a"
    print(
        f"  Sharpe={float(out['sharpe']):.3f}  Calmar={float(out['calmar']):.2f}  "
        f"maxDD%={float(out['max_drawdown_pct']):.2f}  CAGR={_cgs}  endEq=${float(out['end_equity']):,.0f}"
    )
    print(
        f"\n  Re-run merged %% book (no Theta): .venv/bin/python RenTech/strategy_stack/"
        f"portfolio_vrp_plus_vxx.py --allocation-json {args.out_json} --vrp-trades {vrp_path!s}"
    )
    print("=" * 72)


if __name__ == "__main__":
    main()
