#!/usr/bin/env python3
"""
Equal-weight portfolio of selected macro ETF option sleeves (Theta 15:45 monthly rolls).

Default sleeves (from complement scan):
  TLT: buy_write_pmcc, bull_call_spread, butterfly_spread
  USO: butterfly_spread, iron_condor, jade_lizard
  DBC: put_diagonal
  GLD: putw_like (cash-secured monthly put)

Allocation modes (``--allocation``):

- **equal_weight** (default): arithmetic mean of sleeve equity curves — each sleeve
  uses ``--capital`` notional but only 1/N of its PnL counts toward the book.
- **stacked**: ``capital + sum(sleeve_pnl)`` — every sleeve runs at full ``--capital``
  concurrently (~N× notional on one account; boosts return and drawdown).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/macro_aw_options_portfolio.py \\
      --start-date 2016-04-01 --end-date 2026-04-30 \\
      --capital 100000 \\
      --out-prefix RenTech/data/logs/macro_aw_options_portfolio

Optional VRP blend::

    ... macro_aw_options_portfolio.py \\
      --with-vrp-csv RenTech/data/logs/portfolio_opt_10dd_sharpe_fullvrp.csv \\
      --vrp-col eq_vrp_only --vrp-frac 0.50 \\
      --out-prefix RenTech/data/logs/macro_aw_plus_vrp
"""

from __future__ import annotations

import argparse
import importlib.util
import math
import sys
from dataclasses import dataclass
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))


@dataclass(frozen=True)
class SleeveSpec:
    ticker: str
    strategy: str  # benchmark structure name or "putw_like"
    label: str = ""

    @property
    def col(self) -> str:
        return self.label or f"{self.ticker}_{self.strategy}"


# User-selected complement sleeves
DEFAULT_SLEEVES: tuple[SleeveSpec, ...] = (
    SleeveSpec("TLT", "buy_write_pmcc"),
    SleeveSpec("TLT", "bull_call_spread"),
    SleeveSpec("TLT", "butterfly_spread"),
    SleeveSpec("USO", "butterfly_spread"),
    SleeveSpec("USO", "iron_condor"),
    SleeveSpec("USO", "jade_lizard"),
    SleeveSpec("DBC", "put_diagonal"),
    SleeveSpec("GLD", "putw_like"),
)


def _load_module(name: str, path: Path):
    spec = importlib.util.spec_from_file_location(name, path)
    assert spec and spec.loader
    mod = importlib.util.module_from_spec(spec)
    sys.modules[name] = mod
    spec.loader.exec_module(mod)
    return mod


_BENCH = _load_module(
    "bench_opt_by_ticker",
    Path(__file__).with_name("benchmark_option_strategies_by_ticker.py"),
)
_PUTW = _load_module(
    "bench_putw_multi",
    Path(__file__).with_name("benchmark_putw_like_multi_ticker.py"),
)


def _estimate_trade_margin(trades: "pd.DataFrame") -> "pd.Series":
    """Return per-trade margin_reserved_usd estimate (portfolio margin rules, 1 contract).

    Rules (1 contract = 100 multiplier):
    - Short-put strategies (putw_like, put_diagonal, jade_lizard): 20% × spot_entry × 100
    - Spread strategies where risk is bounded (iron_condor, put_credit_spread):
      8% × spot_entry × 100  (≈ spread width assumption)
    - Debit strategies, no additional margin (butterfly_spread, bull_call_spread,
      buy_write_pmcc): 0
    - Unknown / missing spot: flat $500 estimate.
    """
    import numpy as np

    SHORT_PUT_MARGIN_RATE = 0.20   # 20% of underlying (Reg-T short put rule)
    SPREAD_MARGIN_RATE    = 0.08   # ≈ 8% of underlying (spread width proxy)
    FALLBACK_MARGIN       = 500.0  # when spot_entry is zero / missing

    SHORT_PUT_STRATEGIES  = {"putw_like", "put_diagonal", "jade_lizard"}
    SPREAD_STRATEGIES     = {"iron_condor", "put_credit_spread"}
    DEBIT_STRATEGIES      = {"butterfly_spread", "bull_call_spread", "buy_write_pmcc", "long_strangle"}

    margin = pd.Series(np.nan, index=trades.index, dtype=float)
    spot = pd.to_numeric(trades.get("spot_entry", pd.Series(0.0, index=trades.index)), errors="coerce").fillna(0.0)
    strat = trades.get("strategy", pd.Series("", index=trades.index)).fillna("").astype(str)

    for i in trades.index:
        s = strat.loc[i]
        sp = spot.loc[i]
        if s in DEBIT_STRATEGIES:
            margin.loc[i] = 0.0
        elif s in SHORT_PUT_STRATEGIES:
            margin.loc[i] = (sp * SHORT_PUT_MARGIN_RATE * 100.0) if sp > 0 else FALLBACK_MARGIN
        elif s in SPREAD_STRATEGIES:
            margin.loc[i] = (sp * SPREAD_MARGIN_RATE * 100.0) if sp > 0 else FALLBACK_MARGIN
        else:
            margin.loc[i] = FALLBACK_MARGIN
    return margin.fillna(FALLBACK_MARGIN)


def _equity_from_exit_pnl(
    trades_df: pd.DataFrame,
    start: pd.Timestamp,
    end: pd.Timestamp,
    starting_capital: float,
    *,
    date_col: str = "exit_date",
    pnl_col: str = "pnl_usd",
) -> pd.Series:
    idx = pd.bdate_range(start, end)
    running = float(starting_capital)
    vals: list[float] = []
    if trades_df.empty:
        return pd.Series(running, index=idx, dtype=float)
    t = trades_df.copy()
    t[date_col] = pd.to_datetime(t[date_col], errors="coerce").dt.normalize()
    pnl_by = t.groupby(date_col, as_index=True)[pnl_col].sum()
    for d in idx:
        running += float(pnl_by.get(pd.Timestamp(d).normalize(), 0.0))
        vals.append(running)
    return pd.Series(vals, index=idx, dtype=float)


def _run_structure_sleeve(
    spec: SleeveSpec,
    theta_dir: Path,
    start: pd.Timestamp,
    end: pd.Timestamp,
    capital: float,
    *,
    strict_legs: bool,
) -> tuple[pd.DataFrame, pd.Series]:
    opt = _BENCH.load_option_rows(theta_dir, spec.ticker, start, end)
    trades = _BENCH.run_strategy(
        opt,
        spec.ticker,
        spec.strategy,
        start,
        end,
        approximate_missing=not strict_legs,
    )
    tdf = pd.DataFrame([t.__dict__ for t in trades])
    eq = _equity_from_exit_pnl(tdf, start, end, capital)
    return tdf, eq


def _run_putw_sleeve(
    spec: SleeveSpec,
    theta_dir: Path,
    start: pd.Timestamp,
    end: pd.Timestamp,
    capital: float,
    *,
    dte_target: int,
    min_dte: int,
    max_dte: int,
) -> tuple[pd.DataFrame, pd.Series]:
    tdf, _ = _PUTW.run_putw_like(
        theta_dir,
        spec.ticker,
        start,
        end,
        dte_target=dte_target,
        min_dte=min_dte,
        max_dte=max_dte,
        starting_capital=capital,
    )
    eq = _PUTW.build_equity_curve(tdf, start, end, capital)
    return tdf, eq


def _sharpe(eq: pd.Series) -> float:
    r = eq.pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)
    std = float(r.std(ddof=1))
    if std < 1e-12:
        return float("nan")
    return float(r.mean()) / std * math.sqrt(252.0)


def _max_dd_pct(eq: pd.Series) -> float:
    peak = eq.cummax()
    dd = eq / peak - 1.0
    return float(dd.min()) * 100.0


def _metrics(eq: pd.Series, capital: float) -> dict:
    end = float(eq.iloc[-1])
    years = max((eq.index[-1] - eq.index[0]).days / 365.25, 1e-9)
    ret_pct = (end / capital - 1.0) * 100.0
    cagr = ((end / capital) ** (1.0 / years) - 1.0) * 100.0 if end > 0 else float("nan")
    return {
        "ending_capital": end,
        "total_return_pct": ret_pct,
        "cagr_pct": cagr,
        "sharpe": _sharpe(eq),
        "max_dd_pct": _max_dd_pct(eq),
    }


def _spy_returns(start: pd.Timestamp, end: pd.Timestamp) -> pd.Series:
    from RenTech.strategy_stack.data_loader import DataLoader

    spy = DataLoader().fetch_daily("SPY", period="max")
    spy.index = pd.to_datetime(spy.index).tz_localize(None).normalize()
    ret = spy["close"].astype(float).pct_change().fillna(0.0)
    return ret.loc[(ret.index >= start) & (ret.index <= end)]


def _parse_sleeves(raw: str | None) -> tuple[SleeveSpec, ...]:
    if not raw:
        return DEFAULT_SLEEVES
    out: list[SleeveSpec] = []
    for part in raw.split(","):
        part = part.strip()
        if not part:
            continue
        if ":" not in part:
            raise ValueError(f"Expected TICKER:strategy, got {part!r}")
        t, s = part.split(":", 1)
        out.append(SleeveSpec(t.strip().upper(), s.strip()))
    return tuple(out)


def main() -> None:
    ap = argparse.ArgumentParser(
        description="Equal-weight macro AW options portfolio (TLT/USO/DBC/GLD sleeves).",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )
    ap.add_argument("--theta-dir", type=Path, default=Path("RenTech/data/theta_chunks"))
    ap.add_argument("--start-date", type=str, default="2016-04-01")
    ap.add_argument("--end-date", type=str, default="2026-04-30")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument(
        "--sleeves",
        type=str,
        default=None,
        help="Override default list: TICKER:strategy,... (strategy=putw_like for CSP).",
    )
    ap.add_argument("--strict-legs", action="store_true")
    ap.add_argument("--dte-target", type=int, default=30)
    ap.add_argument("--min-dte", type=int, default=20)
    ap.add_argument("--max-dte", type=int, default=45)
    ap.add_argument("--with-vrp-csv", type=Path, default=None, help="Optional VRP daily equity CSV.")
    ap.add_argument("--vrp-col", type=str, default="eq_vrp_only")
    ap.add_argument("--vrp-frac", type=float, default=0.50, help="Weight on VRP when blending.")
    ap.add_argument(
        "--allocation",
        choices=("equal_weight", "stacked"),
        default="equal_weight",
        help="equal_weight = mean of sleeve curves; stacked = sum all sleeve PnL at full capital.",
    )
    ap.add_argument(
        "--out-prefix",
        type=Path,
        default=Path("RenTech/data/logs/macro_aw_options_portfolio"),
    )
    args = ap.parse_args()

    sleeves = _parse_sleeves(args.sleeves)
    start = pd.Timestamp(args.start_date).normalize()
    end = pd.Timestamp(args.end_date).normalize()
    cap = float(args.capital)
    theta_dir = args.theta_dir.expanduser().resolve()
    prefix = args.out_prefix.expanduser()
    prefix.parent.mkdir(parents=True, exist_ok=True)

    print("=" * 92)
    print(f"Macro AW options portfolio  |  {start.date()} → {end.date()}  |  capital ${cap:,.0f}")
    print(f"Sleeves ({len(sleeves)}, {args.allocation}): " + ", ".join(s.col for s in sleeves))
    print("=" * 92)

    equity_cols: dict[str, pd.Series] = {}
    summary_rows: list[dict] = []
    trades_parts: list[pd.DataFrame] = []

    for spec in sleeves:
        print(f"  Running {spec.col} …", flush=True)
        if spec.strategy == "putw_like":
            tdf, eq = _run_putw_sleeve(
                spec, theta_dir, start, end, cap,
                dte_target=int(args.dte_target),
                min_dte=int(args.min_dte),
                max_dte=int(args.max_dte),
            )
        else:
            tdf, eq = _run_structure_sleeve(
                spec, theta_dir, start, end, cap, strict_legs=bool(args.strict_legs),
            )
        equity_cols[spec.col] = eq
        m = _metrics(eq, cap)
        n_trades = int(len(tdf))
        summary_rows.append({
            "sleeve": spec.col,
            "ticker": spec.ticker,
            "strategy": spec.strategy,
            "trades": n_trades,
            **m,
        })
        if not tdf.empty:
            tdf = tdf.copy()
            tdf["sleeve"] = spec.col
            trades_parts.append(tdf)

    curve_df = pd.concat(equity_cols, axis=1)
    if args.allocation == "stacked":
        port_eq = cap + sum(curve_df[c] - cap for c in curve_df.columns)
        port_label = "PORTFOLIO_STACKED"
    else:
        port_eq = curve_df.mean(axis=1)
        port_label = "PORTFOLIO_EQUAL_WEIGHT"
    equity_cols[port_label] = port_eq

    spy_ret = _spy_returns(start, end)
    port_ret = port_eq.pct_change().fillna(0.0)
    corr_spy = float(port_ret.corr(spy_ret.reindex(port_ret.index).fillna(0.0)))

    # Sleeve correlation matrix (daily returns)
    ret_df = curve_df.pct_change().fillna(0.0)
    corr_mat = ret_df.corr()

    alloc_note = "sum of PnL" if args.allocation == "stacked" else "mean of curves"
    print(f"\nPer-sleeve (each at full --capital notional; portfolio = {alloc_note}):")
    print("-" * 92)
    for row in summary_rows:
        print(
            f"  {row['sleeve']:<28}  trades {row['trades']:>4}  "
            f"ret {row['total_return_pct']:>7.1f}%  Sharpe {row['sharpe']:>5.2f}  "
            f"MaxDD {row['max_dd_pct']:>6.2f}%"
        )

    pm = _metrics(port_eq, cap)
    print("-" * 92)
    print(
        f"  {port_label:<28}  "
        f"ret {pm['total_return_pct']:>7.1f}%  CAGR {pm['cagr_pct']:>5.2f}%  "
        f"Sharpe {pm['sharpe']:>5.2f}  MaxDD {pm['max_dd_pct']:>6.2f}%  "
        f"ρ(SPY) {corr_spy:>+.3f}"
    )
    if args.allocation == "stacked":
        print(f"  (Implied book notional ≈ {len(sleeves)} × ${cap:,.0f} = ${len(sleeves) * cap:,.0f})")

    # Optional VRP blend
    blend_eq = None
    if args.with_vrp_csv is not None:
        vrp_path = args.with_vrp_csv.expanduser()
        if not vrp_path.is_file():
            raise SystemExit(f"VRP CSV not found: {vrp_path}")
        vrp_df = pd.read_csv(vrp_path, index_col=0, parse_dates=True)
        vrp_df.index = pd.to_datetime(vrp_df.index).normalize()
        if args.vrp_col not in vrp_df.columns:
            raise SystemExit(f"Column {args.vrp_col!r} not in {vrp_path}")
        vrp_eq = vrp_df[args.vrp_col].dropna().sort_index()
        w_vrp = float(args.vrp_frac)
        if not 0.0 <= w_vrp <= 1.0:
            raise SystemExit("--vrp-frac must be in [0, 1]")
        w_macro = 1.0 - w_vrp
        idx = port_eq.index.intersection(vrp_eq.index).sort_values()
        vrp_al = vrp_eq.reindex(idx).ffill()
        macro_al = port_eq.reindex(idx).ffill()
        blend_eq = cap + w_vrp * (vrp_al - cap) + w_macro * (macro_al - cap)
        equity_cols[f"BLEND_{w_vrp:.0%}_VRP_{w_macro:.0%}_MACRO"] = blend_eq.reindex(port_eq.index).ffill()
        bm = _metrics(blend_eq, cap)
        corr_vrp = float(port_ret.reindex(idx).corr(vrp_al.pct_change().fillna(0.0)))
        print("-" * 92)
        print(
            f"  BLEND {w_vrp:.0%} VRP + {w_macro:.0%} macro  "
            f"ret {bm['total_return_pct']:>7.1f}%  Sharpe {bm['sharpe']:>5.2f}  "
            f"MaxDD {bm['max_dd_pct']:>6.2f}%  ρ(vrp,macro) {corr_vrp:>+.3f}"
        )

    # Write artifacts
    daily_path = Path(f"{prefix}_daily.csv")
    out_daily = pd.concat(equity_cols, axis=1)
    out_daily.index.name = "date"
    out_daily.to_csv(daily_path)

    summary_path = Path(f"{prefix}_sleeves_summary.csv")
    pd.DataFrame(summary_rows).to_csv(summary_path, index=False)

    corr_path = Path(f"{prefix}_sleeve_corr.csv")
    corr_mat.to_csv(corr_path)

    if trades_parts:
        trades_path = Path(f"{prefix}_trades.csv")
        all_trades = pd.concat(trades_parts, ignore_index=True)
        all_trades["margin_reserved_usd"] = _estimate_trade_margin(all_trades)
        all_trades.to_csv(trades_path, index=False)
    else:
        trades_path = None

    meta = {
        "start_date": str(start.date()),
        "end_date": str(end.date()),
        "capital": cap,
        "n_sleeves": len(sleeves),
        "sleeves": [s.col for s in sleeves],
        "allocation": str(args.allocation),
        "portfolio": pm,
        "corr_vs_spy": corr_spy,
        "implied_notional_usd": len(sleeves) * cap if args.allocation == "stacked" else cap,
    }
    if blend_eq is not None:
        meta["vrp_blend_frac"] = float(args.vrp_frac)
        meta["blend"] = _metrics(blend_eq.reindex(port_eq.index).ffill(), cap)

    meta_path = Path(f"{prefix}_meta.json")
    import json

    meta_path.write_text(json.dumps(meta, indent=2), encoding="utf-8")

    print(f"\nWrote daily equity → {daily_path}")
    print(f"Wrote sleeve summary → {summary_path}")
    print(f"Wrote correlation matrix → {corr_path}")
    if trades_path:
        print(f"Wrote trades → {trades_path}")
    print(f"Wrote meta → {meta_path}")


if __name__ == "__main__":
    main()
