#!/usr/bin/env python3
"""
Rank monthly-roll option sleeves on macro ETFs vs SPY (and optional VRP daily PnL).

Uses the same structures as ``benchmark_option_strategies_by_ticker.py``.
Outputs a CSV sorted for **low SPY correlation** then **Sharpe**.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/analyze_macro_option_complement.py \\
      --tickers USO,TLT,IEF,DBC,QQQ,GLD,IWM \\
      --start-date 2016-04-01 --end-date 2026-04-30 \\
      --out-csv RenTech/data/logs/macro_option_complement_ranked.csv
"""

from __future__ import annotations

import argparse
import importlib.util
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))

# Load sibling module without package rename churn.
_BENCH_PATH = Path(__file__).with_name("benchmark_option_strategies_by_ticker.py")
_spec = importlib.util.spec_from_file_location("bench_opt_by_ticker", _BENCH_PATH)
assert _spec and _spec.loader
_bench = importlib.util.module_from_spec(_spec)
sys.modules[_spec.name] = _bench
_spec.loader.exec_module(_bench)

STRATEGIES = [
    "put_diagonal",
    "put_credit_spread",
    "iron_condor",
    "buy_write_pmcc",
    "long_strangle",
    "jade_lizard",
    "butterfly_spread",
    "bull_call_spread",
]


def _daily_pnl_from_trades(trades: list, start: pd.Timestamp, end: pd.Timestamp) -> pd.Series:
    idx = pd.bdate_range(start, end)
    if not trades:
        return pd.Series(0.0, index=idx)
    df = pd.DataFrame([t.__dict__ for t in trades])
    df["exit_date"] = pd.to_datetime(df["exit_date"]).dt.normalize()
    pnl_by = df.groupby("exit_date")["pnl_usd"].sum()
    vals = [float(pnl_by.get(pd.Timestamp(d).normalize(), 0.0)) for d in idx]
    return pd.Series(vals, index=idx, dtype=np.float64)


def _equity_returns(pnl: pd.Series, start_cap: float) -> pd.Series:
    eq = start_cap + pnl.cumsum()
    return eq.pct_change().replace([np.inf, -np.inf], np.nan).fillna(0.0)


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

    dl = DataLoader()
    spy = dl.fetch_daily("SPY", period="max")
    if spy.empty:
        raise RuntimeError("No SPY daily data")
    spy = spy.copy()
    spy.index = pd.to_datetime(spy.index).tz_localize(None).normalize()
    close = spy["close"].astype(np.float64)
    ret = close.pct_change().fillna(0.0)
    ret = ret.loc[(ret.index >= start) & (ret.index <= end)]
    return ret.astype(np.float64)


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

    dl = DataLoader()
    df = dl.fetch_daily(ticker, period="max")
    if df.empty:
        return pd.Series(dtype=np.float64)
    df = df.copy()
    df.index = pd.to_datetime(df.index).tz_localize(None).normalize()
    ret = df["close"].astype(np.float64).pct_change().fillna(0.0)
    return ret.loc[(ret.index >= start) & (ret.index <= end)].astype(np.float64)


def _corr(a: pd.Series, b: pd.Series) -> float:
    idx = a.index.intersection(b.index).sort_values()
    if len(idx) < 30:
        return float("nan")
    x = a.reindex(idx).fillna(0.0)
    y = b.reindex(idx).fillna(0.0)
    if x.std() < 1e-12 or y.std() < 1e-12:
        return float("nan")
    return float(x.corr(y))


def _load_vrp_daily_pnl(path: Path, start: pd.Timestamp, end: pd.Timestamp) -> pd.Series | None:
    if not path.exists():
        return None
    if path.suffix.lower() == ".jsonl":
        rows = []
        import json

        with path.open() as f:
            for line in f:
                line = line.strip()
                if not line:
                    continue
                rows.append(json.loads(line))
        if not rows:
            return None
        key_exit = "exit_date" if "exit_date" in rows[0] else "date"
        key_pnl = "pnl_usd" if "pnl_usd" in rows[0] else "pnl_total"
        df = pd.DataFrame(rows)
        df["d"] = pd.to_datetime(df[key_exit]).dt.normalize()
        pnl = df.groupby("d")[key_pnl].sum()
    else:
        df = pd.read_csv(path)
        dcol = "date" if "date" in df.columns else df.columns[0]
        pcol = next((c for c in ("pnl_usd", "daily_pnl", "pnl") if c in df.columns), None)
        if pcol is None:
            return None
        df["d"] = pd.to_datetime(df[dcol]).dt.normalize()
        pnl = df.groupby("d")[pcol].sum()
    idx = pd.bdate_range(start, end)
    return pd.Series([float(pnl.get(pd.Timestamp(d).normalize(), 0.0)) for d in idx], index=idx)


def main() -> None:
    ap = argparse.ArgumentParser(description="Macro ETF option sleeves vs SPY correlation ranking.")
    ap.add_argument("--theta-dir", type=Path, default=Path("RenTech/data/theta_chunks"))
    ap.add_argument("--tickers", type=str, default="USO,TLT,IEF,DBC,QQQ,GLD,IWM")
    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("--starting-capital", type=float, default=100_000.0)
    ap.add_argument(
        "--vrp-pnl",
        type=Path,
        default=None,
        help="Optional VRP exit-day PnL (JSONL from vrp_backtest_theta --export-trades-jsonl or daily CSV).",
    )
    ap.add_argument("--max-abs-corr-spy", type=float, default=0.35, help="Highlight rows with |ρ| ≤ this.")
    ap.add_argument("--min-trades", type=int, default=80, help="Require at least this many round trips.")
    ap.add_argument("--out-csv", type=Path, default=Path("RenTech/data/logs/macro_option_complement_ranked.csv"))
    args = ap.parse_args()

    tickers = [x.strip().upper() for x in args.tickers.split(",") if x.strip()]
    start = pd.Timestamp(args.start_date).normalize()
    end = pd.Timestamp(args.end_date).normalize()
    spy_ret = _load_spy_returns(start, end)
    vrp_pnl = _load_vrp_daily_pnl(args.vrp_pnl, start, end) if args.vrp_pnl else None
    vrp_ret = None
    if vrp_pnl is not None:
        vrp_ret = _equity_returns(vrp_pnl, float(args.starting_capital))

    rows: list[dict] = []
    for t in tickers:
        opt = _bench.load_option_rows(args.theta_dir, t, start, end)
        und_ret = _load_underlying_returns(t, start, end)
        for s in STRATEGIES:
            trades = _bench.run_strategy(
                opt, t, s, start, end, approximate_missing=True,
            )
            m = _bench.summarize(trades, float(args.starting_capital), start, end)
            pnl = _daily_pnl_from_trades(trades, start, end)
            strat_ret = _equity_returns(pnl, float(args.starting_capital))
            row = {
                "ticker": t,
                "strategy": s,
                "trades": m["trades"],
                "sharpe": m["sharpe"],
                "total_return_pct": m["total_return_pct"],
                "max_dd_pct": m["max_dd_pct"],
                "corr_daily_ret_spy": _corr(strat_ret, spy_ret),
                "corr_daily_ret_underlying": _corr(strat_ret, und_ret) if not und_ret.empty else float("nan"),
            }
            if vrp_ret is not None:
                row["corr_daily_ret_vrp"] = _corr(strat_ret, vrp_ret)
            row["low_corr_spy"] = (
                abs(row["corr_daily_ret_spy"]) <= float(args.max_abs_corr_spy)
                if np.isfinite(row["corr_daily_ret_spy"])
                else False
            )
            row["pass_min_trades"] = int(m["trades"]) >= int(args.min_trades)
            rows.append(row)

    out = pd.DataFrame(rows)
    out["abs_corr_spy"] = out["corr_daily_ret_spy"].abs()
    out = out.sort_values(["abs_corr_spy", "sharpe"], ascending=[True, False], na_position="last")
    args.out_csv.parent.mkdir(parents=True, exist_ok=True)
    out.to_csv(args.out_csv, index=False)

    print("=" * 88)
    print(f"Macro option complement scan  {start.date()} → {end.date()}")
    print(f"Low-|ρ| highlight: |corr vs SPY| ≤ {args.max_abs_corr_spy:.2f}  min_trades ≥ {args.min_trades}")
    if vrp_ret is None:
        print("(No --vrp-pnl: skipping VRP correlation column)")
    print("=" * 88)

    filt = out.loc[out["pass_min_trades"] & out["low_corr_spy"]].head(15)
    cols = ["ticker", "strategy", "trades", "sharpe", "total_return_pct", "max_dd_pct", "corr_daily_ret_spy"]
    if "corr_daily_ret_vrp" in out.columns:
        cols.append("corr_daily_ret_vrp")
    print("\nTop candidates (low SPY ρ, min trades):")
    print(filt[cols].to_string(index=False, float_format=lambda x: f"{x:,.3f}"))

    print(f"\nWrote: {args.out_csv}")


if __name__ == "__main__":
    main()
