#!/usr/bin/env python3
"""
Portfolio weight optimizer for three independent VXX strategies:
  A) Bear call credit spread   (steady theta grinder)
  B) Deep ITM put              (directional decay capture)
  C) Long OTM call             (tail hedge on VIX spikes)

Builds daily equity curves for each, sweeps a 3-D weight grid
(wA + wB + wC = 1), and finds the combination that maximises
risk-adjusted return (Sharpe, Calmar, or total PnL).

Example::

    python RenTech/strategy_stack/optimize_vxx_portfolio.py \
        --contango-mode futures --contango-threshold 0.03 --step-pct 5
"""
from __future__ import annotations

import math, json, sys, shutil
from datetime import datetime, timezone
from pathlib import Path
from dataclasses import asdict

import numpy as np
import pandas as pd

_REPO = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(_REPO))
sys.path.insert(0, str(_REPO / "RenTech" / "strategy_stack"))

from explore_vxx_decay_strategies import (
    _load_contango, _load_chain, _spot_from_chain, _pick_expiry,
    _build_bear_call_credit, _exit_bear_call_from_chain,
    _build_deep_itm_put, _exit_deep_put,
    _build_long_call, _exit_long_call,
    _slipped, MULT, SLIPPAGE, Trade,
    broker_risk_usd_from_built,
    leg_strikes_from_vxx_built,
    resolve_vxx_contracts_and_broker_risk,
    vxx_built_snapshot,
)


# ── run a single strategy → (trades, daily equity Series) ──────────────────

def _run_strategy(
    name: str,
    builder,
    exiter,
    max_loss_field: str | None,
    ct: pd.DataFrame,
    dates: list[pd.Timestamp],
    *,
    contango_mode: str,
    contango_threshold: float,
    vix3m_threshold: float,
    dte_min: int,
    dte_max: int,
    hold_days: int,
    rebalance_every: int,
    stop_loss_mult: float,
    contracts: int | None = None,
    target_broker_risk_usd: float | None = None,
    broker_risk_pct_of_portfolio: float | None = None,
    initial_portfolio_capital: float = 100_000.0,
) -> tuple[list[Trade], pd.Series]:
    trades: list[Trade] = []
    pending = None
    days_held = 0
    equity = 0.0
    daily_eq = {}

    for step, d in enumerate(dates):
        d = pd.Timestamp(d).normalize()
        ct_row = ct.loc[d]
        v3v = float(ct_row.get("vix3m_vix_ratio", np.nan))
        cr_val = float(ct_row.get("contango_ratio_ffill", np.nan))

        # ── exit pass ──
        if pending is not None:
            days_held += 1
            p = pending
            dte_left = int((p["expiration"] - d).days)
            should_exit = (days_held >= p["hold_target"]) or (dte_left <= 1)

            if not should_exit and max_loss_field and p.get("max_loss") is not None:
                chain = _load_chain(d)
                spot_now = _spot_from_chain(chain) if not chain.empty else None
                if spot_now is not None and not chain.empty:
                    pnl_one = exiter(chain, spot_now, p["legs"], p["expiration"])
                    pnl_now = float(pnl_one) * int(p["contracts"])
                    if pnl_now <= -float(p["max_loss"]) * stop_loss_mult:
                        should_exit = True

            if should_exit:
                chain = _load_chain(d)
                spot_now = _spot_from_chain(chain) if not chain.empty else None
                if spot_now is None:
                    spot_now = p["vxx_entry"]
                pnl_one = exiter(chain, spot_now, p["legs"], p["expiration"])
                pnl = float(pnl_one) * int(p["contracts"])
                reason = "time" if days_held >= p["hold_target"] or dte_left <= 1 else "stop_loss"
                equity += pnl
                exp_s = str(pd.Timestamp(p["expiration"]).date())
                ss, ls, ps, cs = leg_strikes_from_vxx_built(p["legs"])
                ml_one = p["legs"].get("max_loss")
                ml_one_f = float(ml_one) if ml_one is not None and math.isfinite(float(ml_one)) else None
                trades.append(
                    Trade(
                        strategy=name,
                        entry_date=str(p["entry_date"]),
                        exit_date=str(d.date()),
                        exit_reason=reason,
                        vxx_entry=p["vxx_entry"],
                        vxx_exit=spot_now,
                        entry_credit_or_debit=float(p["entry_val"]),
                        exit_value=pnl + float(p["entry_val"]),
                        pnl_total=pnl,
                        contango_ratio=p.get("cr", 0.0),
                        vix3m_vix=p.get("v3v", 0.0),
                        broker_risk_usd=float(p["broker_risk_usd"]),
                        contracts=int(p["contracts"]),
                        broker_risk_per_contract_usd=float(p["broker_risk_per_contract"]),
                        underlying="VXX",
                        expiration=exp_s,
                        short_strike=ss,
                        long_strike=ls,
                        put_strike=ps,
                        call_strike=cs,
                        nav_at_entry_usd=float(p.get("nav_at_entry", 0.0)),
                        risk_pct_of_portfolio=float(p.get("risk_pct_of_portfolio", 0.0)),
                        dte_at_entry=int(p.get("dte_at_entry", 0)),
                        dte_at_exit=int(dte_left),
                        hold_target_days=int(p["hold_target"]),
                        days_held=int(days_held),
                        max_loss_one_usd=ml_one_f,
                        max_loss_total_usd=float(p["max_loss"]) if p.get("max_loss") is not None else None,
                        target_broker_risk_usd=p.get("target_broker_risk_usd"),
                        contracts_requested=p.get("contracts_requested"),
                        entry_legs_json=str(p.get("entry_legs_json", "{}")),
                    )
                )
                pending = None
                days_held = 0

        daily_eq[d] = equity

        # ── entry pass ──
        if step % rebalance_every != 0:
            continue
        if pending is not None:
            continue

        in_contango = False
        if contango_mode == "futures":
            in_contango = math.isfinite(cr_val) and cr_val >= contango_threshold
        elif contango_mode == "vix3m":
            in_contango = math.isfinite(v3v) and v3v >= vix3m_threshold
        elif contango_mode == "both":
            in_contango = (math.isfinite(cr_val) and cr_val >= contango_threshold
                           and math.isfinite(v3v) and v3v >= vix3m_threshold)
        if not in_contango:
            continue

        chain = _load_chain(d)
        if chain.empty:
            continue
        spot = _spot_from_chain(chain)
        if spot is None:
            continue
        exp = _pick_expiry(chain, dte_min, dte_max)
        if exp is None:
            continue

        built = builder(chain, spot, exp)
        if built is None:
            continue

        entry_val_one = float(
            built.get("credit", built.get("net_entry", -float(built.get("debit", 0.0))))
        )
        nav_at = max(float(initial_portfolio_capital) + float(equity), 1.0)
        pct = broker_risk_pct_of_portfolio
        tgt_for_row: float | None = None
        if pct is not None:
            if not (0.0 < float(pct) <= 1.0):
                raise ValueError("broker_risk_pct_of_portfolio must be in (0, 1]")
            tgt = max(nav_at * float(pct), 1.0)
            tgt_for_row = float(tgt)
            n_c, per_risk, br_total = resolve_vxx_contracts_and_broker_risk(
                strategy=name,
                built=built,
                entry_val=entry_val_one,
                contracts=None,
                target_broker_risk_usd=tgt,
            )
            pct_applied = float(pct)
        else:
            n_c, per_risk, br_total = resolve_vxx_contracts_and_broker_risk(
                strategy=name,
                built=built,
                entry_val=entry_val_one,
                contracts=contracts,
                target_broker_risk_usd=target_broker_risk_usd,
            )
            pct_applied = 0.0
            if target_broker_risk_usd is not None:
                tgt_for_row = float(target_broker_risk_usd)
        entry_val = entry_val_one * float(n_c)
        ml_one = built.get("max_loss")
        max_loss_total = float(ml_one) * float(n_c) if ml_one is not None else None
        ml_one_f = float(ml_one) if ml_one is not None and math.isfinite(float(ml_one)) else None

        days_to_exp = sum(1 for dd in dates if d < dd <= exp) - 1
        ht = min(hold_days, max(days_to_exp, 1))
        dte_entry = int((exp - d).days)

        pending = {
            "entry_date": d.date(),
            "expiration": exp,
            "legs": built,
            "vxx_entry": spot,
            "entry_val": entry_val,
            "hold_target": ht,
            "max_loss": max_loss_total,
            "broker_risk_usd": br_total,
            "broker_risk_per_contract": per_risk,
            "contracts": n_c,
            "nav_at_entry": float(nav_at),
            "risk_pct_of_portfolio": float(pct_applied),
            "dte_at_entry": int(dte_entry),
            "target_broker_risk_usd": tgt_for_row,
            "contracts_requested": int(contracts) if contracts is not None else None,
            "max_loss_one_usd": ml_one_f,
            "entry_legs_json": json.dumps(vxx_built_snapshot(built)),
            "cr": cr_val if math.isfinite(cr_val) else 0.0,
            "v3v": v3v if math.isfinite(v3v) else 0.0,
        }
        days_held = 0

    eq_series = pd.Series(daily_eq, name=name).sort_index()
    return trades, eq_series


# ── portfolio metrics on a daily equity curve ───────────────────────────────

def _metrics(eq: pd.Series) -> dict:
    if eq.empty or (eq.iloc[-1] == 0 and eq.iloc[0] == 0):
        return {"total_pnl": 0, "sharpe": 0, "max_dd": 0, "calmar": 0,
                "avg_daily": 0, "vol_daily": 0}
    daily_ret = eq.diff().fillna(0)
    total = float(eq.iloc[-1])
    mn = float(daily_ret.mean())
    sd = float(daily_ret.std()) if daily_ret.std() > 0 else 1e-9
    sharpe = (mn / sd) * np.sqrt(252)
    peak = eq.cummax()
    dd = eq - peak
    max_dd = float(dd.min())
    calmar = total / abs(max_dd) if max_dd != 0 else 0
    return {
        "total_pnl": round(total, 1),
        "sharpe": round(sharpe, 3),
        "max_dd": round(max_dd, 1),
        "calmar": round(calmar, 3),
        "avg_daily": round(mn, 3),
        "vol_daily": round(sd, 3),
    }


# ── main ────────────────────────────────────────────────────────────────────

def main():
    import argparse
    ap = argparse.ArgumentParser()
    ap.add_argument("--start", default="2018-06-01")
    ap.add_argument("--end", default="2025-12-31")
    ap.add_argument("--contango-mode", choices=["futures", "vix3m", "both"], default="futures")
    ap.add_argument("--contango-threshold", type=float, default=0.03)
    ap.add_argument("--vix3m-threshold", type=float, default=1.05)
    ap.add_argument("--hold-days", type=int, default=20)
    ap.add_argument("--rebalance-every", type=int, default=10)
    ap.add_argument("--dte-min", type=int, default=21)
    ap.add_argument("--dte-max", type=int, default=45)
    ap.add_argument("--step-pct", type=int, default=5,
                    help="Weight grid step in %% (default 5 → 231 combos)")
    ap.add_argument(
        "--contracts",
        type=int,
        default=None,
        help="VXX option contracts per leg for all three strategies (default 1).",
    )
    ap.add_argument(
        "--target-broker-risk-usd",
        type=float,
        default=None,
        metavar="USD",
        help="Fixed per-entry broker-risk budget; contracts = floor(target / per-contract risk).",
    )
    ap.add_argument(
        "--portfolio-capital",
        type=float,
        default=100_000.0,
        help="Starting NAV anchor for %% sizing (default 100k).",
    )
    ap.add_argument(
        "--broker-risk-pct-of-portfolio",
        type=float,
        default=None,
        metavar="FRAC",
        help="e.g. 0.025 = 2.5%% of NAV before each entry. NAV = --portfolio-capital + this strategy's "
        "cumulative realized PnL so far. Implies contracts; overrides --contracts and --target-broker-risk-usd.",
    )
    ap.add_argument(
        "--trade-audit-dir",
        type=Path,
        default=None,
        help="If set, also writes a run_manifest.json plus copies of the three JSONL trade logs into this directory.",
    )
    ap.add_argument(
        "--trade-audit-tag",
        type=str,
        default="",
        help="Optional subdirectory name under --trade-audit-dir (safe identifier).",
    )
    args = ap.parse_args()

    if args.broker_risk_pct_of_portfolio is not None and (
        args.contracts is not None or args.target_broker_risk_usd is not None
    ):
        print(
            "ERROR: --broker-risk-pct-of-portfolio cannot be combined with --contracts or --target-broker-risk-usd",
            file=sys.stderr,
        )
        sys.exit(1)

    ct = _load_contango()
    all_dates = sorted(ct.index)
    dates = [d for d in all_dates if args.start <= str(d.date()) <= args.end]
    print(f"Date range: {dates[0].date()} → {dates[-1].date()}  ({len(dates)} days)")
    print(f"Contango: {args.contango_mode} >= {args.contango_threshold}")

    common = dict(
        ct=ct, dates=dates,
        contango_mode=args.contango_mode,
        contango_threshold=args.contango_threshold,
        vix3m_threshold=args.vix3m_threshold,
        dte_min=args.dte_min, dte_max=args.dte_max,
        hold_days=args.hold_days,
        rebalance_every=args.rebalance_every,
        stop_loss_mult=2.0,
        contracts=args.contracts,
        target_broker_risk_usd=args.target_broker_risk_usd,
        broker_risk_pct_of_portfolio=args.broker_risk_pct_of_portfolio,
        initial_portfolio_capital=float(args.portfolio_capital),
    )

    # ── Strategy A: bear call credit spread (5% OTM, 15% width) ──
    print("\nRunning Strategy A: bear_call_5otm_15w …", flush=True)
    trades_a, eq_a = _run_strategy(
        "bear_call", lambda c, s, e: _build_bear_call_credit(c, s, e, 0.15, 1.05),
        _exit_bear_call_from_chain, "max_loss", **common,
    )

    # ── Strategy B: deep ITM put (20% ITM) ──
    print("Running Strategy B: deep_itm_put …", flush=True)
    trades_b, eq_b = _run_strategy(
        "deep_put", lambda c, s, e: _build_deep_itm_put(c, s, e, 0.20),
        _exit_deep_put, None, **common,
    )

    # ── Strategy C: long OTM call (10% OTM) ──
    print("Running Strategy C: long_otm_call …", flush=True)
    trades_c, eq_c = _run_strategy(
        "long_call", lambda c, s, e: _build_long_call(c, s, e, 0.10),
        _exit_long_call, None, **common,
    )

    # Align indices
    idx = eq_a.index.union(eq_b.index).union(eq_c.index).sort_values()
    eq_a = eq_a.reindex(idx).ffill().fillna(0)
    eq_b = eq_b.reindex(idx).ffill().fillna(0)
    eq_c = eq_c.reindex(idx).ffill().fillna(0)

    ma, mb, mc = _metrics(eq_a), _metrics(eq_b), _metrics(eq_c)
    print(f"\n{'─'*75}")
    print(f"Strategy A (bear_call):  PnL={ma['total_pnl']:>8}  Sharpe={ma['sharpe']:.2f}  MaxDD={ma['max_dd']:>8}  Calmar={ma['calmar']:.2f}")
    print(f"Strategy B (deep_put):   PnL={mb['total_pnl']:>8}  Sharpe={mb['sharpe']:.2f}  MaxDD={mb['max_dd']:>8}  Calmar={mb['calmar']:.2f}")
    print(f"Strategy C (long_call):  PnL={mc['total_pnl']:>8}  Sharpe={mc['sharpe']:.2f}  MaxDD={mc['max_dd']:>8}  Calmar={mc['calmar']:.2f}")

    # ── correlation matrix ──
    da = eq_a.diff().fillna(0)
    db = eq_b.diff().fillna(0)
    dc = eq_c.diff().fillna(0)
    corr_df = pd.DataFrame({"A_bear": da, "B_dput": db, "C_lcall": dc}).corr()
    print(f"\nDaily return correlations:")
    print(corr_df.to_string(float_format=lambda x: f"{x:+.3f}"))

    # ── 3-D weight sweep (wA + wB + wC = 1) ──
    step = args.step_pct / 100.0
    ticks = np.arange(0, 1 + step / 2, step)
    results = []

    for wa in ticks:
        for wb in ticks:
            wc = round(1.0 - wa - wb, 4)
            if wc < -1e-9 or wc > 1 + 1e-9:
                continue
            wc = max(0.0, wc)
            combo = wa * eq_a + wb * eq_b + wc * eq_c
            m = _metrics(combo)
            m["wA_bear"] = round(wa, 2)
            m["wB_dput"] = round(wb, 2)
            m["wC_lcall"] = round(wc, 2)
            results.append(m)

    df = pd.DataFrame(results)
    n_combos = len(df)

    # ── find optima ──
    best_sharpe_idx = df["sharpe"].idxmax()
    best_calmar_idx = df["calmar"].idxmax()
    best_pnl_idx = df["total_pnl"].idxmax()

    # ── top-N tables ──
    top_n = 20

    print(f"\n{'='*100}")
    print(f"3-WAY WEIGHT SWEEP  ({n_combos} combinations, step={args.step_pct}%)")
    print(f"{'='*100}")

    header = f"{'wA':>5} {'wB':>5} {'wC':>5} {'TotalPnL':>9} {'Sharpe':>7} {'MaxDD':>8} {'Calmar':>7} {'AvgDay':>7} {'VolDay':>7}"

    # Top by Sharpe
    print(f"\n── TOP {top_n} BY SHARPE ──")
    print(header)
    print("─" * 100)
    top_sharpe = df.nlargest(top_n, "sharpe")
    for _, row in top_sharpe.iterrows():
        flags = ""
        if _ == best_sharpe_idx:
            flags += " ★SHARPE"
        if _ == best_calmar_idx:
            flags += " ★CALMAR"
        if _ == best_pnl_idx:
            flags += " ★PNL"
        print(
            f"{row['wA_bear']:>4.0%} {row['wB_dput']:>4.0%} {row['wC_lcall']:>4.0%} "
            f"{row['total_pnl']:>9.0f} {row['sharpe']:>7.2f} "
            f"{row['max_dd']:>8.0f} {row['calmar']:>7.2f} "
            f"{row['avg_daily']:>7.3f} {row['vol_daily']:>7.3f}{flags}"
        )

    # Top by Calmar
    print(f"\n── TOP {top_n} BY CALMAR ──")
    print(header)
    print("─" * 100)
    top_calmar = df.nlargest(top_n, "calmar")
    for _, row in top_calmar.iterrows():
        flags = ""
        if _ == best_sharpe_idx:
            flags += " ★SHARPE"
        if _ == best_calmar_idx:
            flags += " ★CALMAR"
        if _ == best_pnl_idx:
            flags += " ★PNL"
        print(
            f"{row['wA_bear']:>4.0%} {row['wB_dput']:>4.0%} {row['wC_lcall']:>4.0%} "
            f"{row['total_pnl']:>9.0f} {row['sharpe']:>7.2f} "
            f"{row['max_dd']:>8.0f} {row['calmar']:>7.2f} "
            f"{row['avg_daily']:>7.3f} {row['vol_daily']:>7.3f}{flags}"
        )

    # Top by PnL
    print(f"\n── TOP {top_n} BY TOTAL PNL ──")
    print(header)
    print("─" * 100)
    top_pnl = df.nlargest(top_n, "total_pnl")
    for _, row in top_pnl.iterrows():
        flags = ""
        if _ == best_sharpe_idx:
            flags += " ★SHARPE"
        if _ == best_calmar_idx:
            flags += " ★CALMAR"
        if _ == best_pnl_idx:
            flags += " ★PNL"
        print(
            f"{row['wA_bear']:>4.0%} {row['wB_dput']:>4.0%} {row['wC_lcall']:>4.0%} "
            f"{row['total_pnl']:>9.0f} {row['sharpe']:>7.2f} "
            f"{row['max_dd']:>8.0f} {row['calmar']:>7.2f} "
            f"{row['avg_daily']:>7.3f} {row['vol_daily']:>7.3f}{flags}"
        )

    print(f"\n{'='*100}")
    bs = df.iloc[best_sharpe_idx]
    bc = df.iloc[best_calmar_idx]
    bp = df.iloc[best_pnl_idx]
    print(f"★ MAX SHARPE:  A={bs['wA_bear']:.0%}  B={bs['wB_dput']:.0%}  C={bs['wC_lcall']:.0%}  →  "
          f"Sharpe {bs['sharpe']:.2f}  PnL ${bs['total_pnl']:.0f}  MaxDD ${bs['max_dd']:.0f}  Calmar {bs['calmar']:.2f}")
    print(f"★ MAX CALMAR:  A={bc['wA_bear']:.0%}  B={bc['wB_dput']:.0%}  C={bc['wC_lcall']:.0%}  →  "
          f"Calmar {bc['calmar']:.2f}  PnL ${bc['total_pnl']:.0f}  MaxDD ${bc['max_dd']:.0f}  Sharpe {bc['sharpe']:.2f}")
    print(f"★ MAX PNL:     A={bp['wA_bear']:.0%}  B={bp['wB_dput']:.0%}  C={bp['wC_lcall']:.0%}  →  "
          f"PnL ${bp['total_pnl']:.0f}  Sharpe {bp['sharpe']:.2f}  MaxDD ${bp['max_dd']:.0f}  Calmar {bp['calmar']:.2f}")
    print(f"{'='*100}")

    # ── save outputs ──
    out_dir = Path(_REPO) / "RenTech" / "data" / "logs"
    out_dir.mkdir(parents=True, exist_ok=True)
    for label, tlist in [("bear_call", trades_a), ("deep_put", trades_b), ("long_call", trades_c)]:
        p = out_dir / f"vxx_portfolio_{label}.jsonl"
        with p.open("w") as f:
            for t in tlist:
                f.write(json.dumps(asdict(t)) + "\n")
        print(f"Saved {len(tlist)} trades → {p}")

    sweep_path = out_dir / "vxx_portfolio_3way_sweep.csv"
    df.to_csv(sweep_path, index=False)
    print(f"Weight sweep ({n_combos} combos) → {sweep_path}")

    if args.trade_audit_dir is not None:
        audit_root = args.trade_audit_dir.expanduser().resolve()
        tag = (args.trade_audit_tag or "").strip()
        audit_dir = audit_root / tag if tag else audit_root
        audit_dir.mkdir(parents=True, exist_ok=True)
        manifest = {
            "created_utc": datetime.now(timezone.utc).isoformat(),
            "script": "optimize_vxx_portfolio.py",
            "args": {
                "start": args.start,
                "end": args.end,
                "contango_mode": args.contango_mode,
                "contango_threshold": float(args.contango_threshold),
                "vix3m_threshold": float(args.vix3m_threshold),
                "hold_days": int(args.hold_days),
                "rebalance_every": int(args.rebalance_every),
                "dte_min": int(args.dte_min),
                "dte_max": int(args.dte_max),
                "step_pct": int(args.step_pct),
                "contracts": args.contracts,
                "target_broker_risk_usd": args.target_broker_risk_usd,
                "portfolio_capital": float(args.portfolio_capital),
                "broker_risk_pct_of_portfolio": args.broker_risk_pct_of_portfolio,
            },
            "outputs": {
                "default_logs_dir": str(out_dir),
                "sweep_csv": str(sweep_path),
            },
            "trade_counts": {
                "bear_call": len(trades_a),
                "deep_put": len(trades_b),
                "long_call": len(trades_c),
            },
        }
        (audit_dir / "run_manifest.json").write_text(json.dumps(manifest, indent=2), encoding="utf-8")
        for label in ("bear_call", "deep_put", "long_call"):
            src = out_dir / f"vxx_portfolio_{label}.jsonl"
            dst = audit_dir / f"vxx_portfolio_{label}.jsonl"
            shutil.copy2(src, dst)
        print(f"Trade audit bundle → {audit_dir}")


if __name__ == "__main__":
    main()
