#!/usr/bin/env python3
"""
Long **VXX puts** on **sweet-spot** term structure: rich M1→M2 roll + VX1 elevated
but not in a blow-off (0–5% above 60d SMA).  Sweeps **moneyness** and **expiration**.

Signals
-------
* ``roll_cost_m1_m2`` = VX2/VX1 − 1  (≥ 7% default)
* ``vx1_elev_vs_ma`` = VX1 / SMA60 − 1  (0%–5% sweet band)
* ``vxx_vs_ma`` optional cap (VXX not already extended vs its MA)

Example::

    cd /Users/robzingale/trading_bot
    .venv/bin/python RenTech/data_pipeline/download_cboe_vix_futures.py
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/backtest_vxx_long_put_roll_carry.py \\
        --start 2016-01-01 --end 2026-12-31 --sweet-spot-sweep --preload-chains
"""

from __future__ import annotations

import argparse
import itertools
import json
import math
import sys
from dataclasses import asdict, 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))

from RenTech.strategy_stack.explore_vxx_decay_strategies import (
    Trade,
    _exit_deep_put,
    _load_contango,
    _nearest_strike,
    _pick_expiry,
    _slipped,
    _spot_from_chain,
    resolve_vxx_contracts_and_broker_risk,
    vxx_built_snapshot,
)
from RenTech.strategy_stack.analyze_vxx_forward_returns import load_vxx_close
from RenTech.strategy_stack.backtest_vxx_vx1_vx3_strategies import (
    START_CAPITAL,
    _load_chain_cached,
    daily_equity_metrics,
)

DATA_DIR = _REPO / "RenTech" / "data"
LOGS = DATA_DIR / "logs"
CONTANGO_PATH = DATA_DIR / "vix_futures_cboe.parquet"

# Sweet-spot defaults from forward-return analysis
SWEET_MIN_ROLL = 0.07
SWEET_MAX_ROLL = 0.12
SWEET_MIN_ELEV = 0.0
SWEET_MAX_ELEV = 0.05
SWEET_MAX_VXX_VS_MA = 0.05


@dataclass
class EntryParams:
    min_roll: float = SWEET_MIN_ROLL
    max_roll: float | None = SWEET_MAX_ROLL
    min_elev_vs_ma: float = SWEET_MIN_ELEV
    max_elev_vs_ma: float = SWEET_MAX_ELEV
    max_vxx_vs_ma: float | None = SWEET_MAX_VXX_VS_MA
    ma_window: int = 60


def enrich_vix_panel(ct: pd.DataFrame, start: str, end: str, ma_window: int) -> pd.DataFrame:
    ct = ct.copy()
    ct.index = pd.to_datetime(ct.index).normalize()
    s1 = ct["vx1_settle"].fillna(ct.get("vx1_close"))
    s2 = ct["vx2_settle"].fillna(ct.get("vx2_close"))
    ct["vx1"] = s1
    ct["roll_cost_m1_m2"] = (s2 / s1) - 1.0
    min_p = max(ma_window // 2, 20)
    ct["vx1_sma"] = ct["vx1"].rolling(ma_window, min_periods=min_p).mean()
    ct["vx1_elev_vs_ma"] = (ct["vx1"] / ct["vx1_sma"]) - 1.0
    ct["roll_cost_m1_m2_ffill"] = ct["roll_cost_m1_m2"].ffill()
    ct["vx1_elev_vs_ma_ffill"] = ct["vx1_elev_vs_ma"].ffill()

    vxx = load_vxx_close(start, end)
    vxx_sma = vxx.rolling(ma_window, min_periods=min_p).mean()
    ct["vxx_close"] = vxx.reindex(ct.index)
    ct["vxx_vs_ma"] = (ct["vxx_close"] / vxx_sma.reindex(ct.index)) - 1.0
    ct["vxx_vs_ma_ffill"] = ct["vxx_vs_ma"].ffill()
    return ct


def _build_long_put(
    chain: pd.DataFrame,
    spot: float,
    exp: pd.Timestamp,
    moneyness: float,
) -> dict | None:
    """
  Buy put at strike ≈ spot × (1 + moneyness).
  moneyness > 0 → ITM, 0 → ATM, < 0 → OTM.
    """
    target_k = spot * (1.0 + moneyness)
    row = _nearest_strike(chain, target_k, "P", exp)
    if row is None:
        return None
    k = float(row["strike"])
    debit = _slipped(float(row["mid"]), "buy") * 100.0
    if debit <= 0:
        return None
    label = "atm"
    if moneyness > 0.01:
        label = "itm"
    elif moneyness < -0.01:
        label = "otm"
    return {
        "debit": debit,
        "strike": k,
        "right": "P",
        "moneyness": moneyness,
        "moneyness_label": label,
    }


def entry_signal(row: pd.Series, p: EntryParams) -> bool:
    roll = float(row.get("roll_cost_m1_m2_ffill", np.nan))
    elev = float(row.get("vx1_elev_vs_ma_ffill", np.nan))
    if not (math.isfinite(roll) and math.isfinite(elev)):
        return False
    if roll < p.min_roll:
        return False
    if p.max_roll is not None and roll >= p.max_roll:
        return False
    if elev < p.min_elev_vs_ma or elev > p.max_elev_vs_ma:
        return False
    if p.max_vxx_vs_ma is not None:
        vxx_e = float(row.get("vxx_vs_ma_ffill", row.get("vxx_vs_ma", np.nan)))
        if math.isfinite(vxx_e) and vxx_e > p.max_vxx_vs_ma:
            return False
    return True


def run_long_put_roll_carry(
    ct: pd.DataFrame,
    dates: list[pd.Timestamp],
    *,
    entry: EntryParams,
    moneyness: float = 0.05,
    dte_min: int = 21,
    dte_max: int = 45,
    hold_days: int = 15,
    rebalance_every: int = 5,
    risk_budget_usd: float = 3000.0,
    profit_take_frac: float | None = 0.5,
    stop_loss_mult: float = 1.0,
) -> list[Trade]:
    trades: list[Trade] = []
    pending: dict | None = None
    days_held = 0

    for step, d in enumerate(dates):
        d = pd.Timestamp(d).normalize()
        if d not in ct.index:
            continue
        row = ct.loc[d]
        roll = float(row.get("roll_cost_m1_m2_ffill", np.nan))
        elev = float(row.get("vx1_elev_vs_ma_ffill", np.nan))

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

            chain = _load_chain_cached(d)
            spot_now = _spot_from_chain(chain) if not chain.empty else p["vxx_entry"]
            pnl_one = _exit_deep_put(chain, spot_now, p["legs"], exp)
            pnl = float(pnl_one) * int(p["contracts"]) * float(p["scale"])

            if not should_exit and profit_take_frac is not None and p.get("profit_target_usd"):
                if pnl >= float(p["profit_target_usd"]) * float(profit_take_frac):
                    should_exit = True
            if not should_exit and p.get("max_loss_usd"):
                if pnl <= -float(p["max_loss_usd"]) * stop_loss_mult:
                    should_exit = True

            if should_exit:
                trades.append(
                    Trade(
                        strategy="long_put_roll_carry",
                        entry_date=str(p["entry_date"]),
                        exit_date=str(d.date()),
                        exit_reason="time",
                        vxx_entry=p["vxx_entry"],
                        vxx_exit=spot_now,
                        entry_credit_or_debit=-float(p["entry_debit"]),
                        exit_value=pnl + float(p["entry_debit"]),
                        pnl_total=pnl,
                        contango_ratio=roll if math.isfinite(roll) else 0.0,
                        vix3m_vix=elev if math.isfinite(elev) else 0.0,
                        broker_risk_usd=float(p["broker_risk_usd"]),
                        contracts=int(p["contracts"]),
                        put_strike=float(p["legs"]["strike"]),
                        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),
                        entry_legs_json=json.dumps(vxx_built_snapshot(p["legs"])),
                    )
                )
                pending = None
                days_held = 0

        if step % rebalance_every != 0 or pending is not None:
            continue
        if not entry_signal(row, entry):
            continue

        chain = _load_chain_cached(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 = _build_long_put(chain, spot, exp, moneyness)
        if built is None:
            continue

        debit_one = float(built["debit"])
        n_c, _, br_tot = resolve_vxx_contracts_and_broker_risk(
            strategy="long_put_roll_carry",
            built=built,
            entry_val=debit_one,
            contracts=None,
            target_broker_risk_usd=risk_budget_usd,
        )
        scale = risk_budget_usd / max(br_tot, 1.0)
        days_to_exp = sum(1 for dd in dates if d < dd <= exp) - 1
        ht = min(hold_days, max(days_to_exp, 1))

        pending = {
            "entry_date": d.date(),
            "expiration": exp,
            "legs": built,
            "vxx_entry": spot,
            "entry_debit": debit_one * n_c * scale,
            "hold_target": ht,
            "max_loss_usd": debit_one * n_c * scale,
            "profit_target_usd": debit_one * n_c * scale * 2.0,
            "broker_risk_usd": br_tot * scale,
            "contracts": n_c,
            "scale": scale,
            "dte_at_entry": int((exp - d).days),
        }
        days_held = 0

    return trades


def _trade_enrichment(trades: list[Trade]) -> dict:
    if not trades:
        return {}
    pnls = [t.pnl_total for t in trades]
    spot_rets = []
    for t in trades:
        if t.vxx_entry and t.vxx_entry > 0:
            spot_rets.append((t.vxx_exit - t.vxx_entry) / t.vxx_entry * 100)
    return {
        "avg_pnl_usd": round(float(np.mean(pnls)), 2),
        "avg_vxx_spot_ret_pct": round(float(np.mean(spot_rets)), 3) if spot_rets else None,
        "median_dte_entry": int(np.median([t.dte_at_entry for t in trades if t.dte_at_entry])),
    }


def run_option_sweep(
    ct: pd.DataFrame,
    dates: list[pd.Timestamp],
    *,
    entry: EntryParams,
    capital: float,
    risk_budget: float,
) -> pd.DataFrame:
    """Sweep moneyness × DTE window × hold days."""
    moneyness_grid = [-0.03, 0.0, 0.02, 0.03, 0.05, 0.07, 0.10]
    dte_windows = [(14, 28), (21, 35), (21, 45), (28, 50), (35, 60)]
    hold_days_grid = [10, 15, 20]

    rows = []
    total = len(moneyness_grid) * len(dte_windows) * len(hold_days_grid)
    done = 0
    for mny, (dte_lo, dte_hi), hold in itertools.product(
        moneyness_grid, dte_windows, hold_days_grid
    ):
        done += 1
        if done % 10 == 0:
            print(f"  sweep [{done}/{total}] mny={mny:+.0%} dte={dte_lo}-{dte_hi} hold={hold}", flush=True)
        tr = run_long_put_roll_carry(
            ct,
            dates,
            entry=entry,
            moneyness=mny,
            dte_min=dte_lo,
            dte_max=dte_hi,
            hold_days=hold,
            risk_budget_usd=risk_budget,
        )
        m = daily_equity_metrics(tr, dates, capital)
        extra = _trade_enrichment(tr)
        label = "ATM" if abs(mny) < 0.01 else ("ITM" if mny > 0 else "OTM")
        rows.append({
            "moneyness": mny,
            "moneyness_label": label,
            "dte_min": dte_lo,
            "dte_max": dte_hi,
            "hold_days": hold,
            "sharpe": m.get("sharpe", 0),
            "n": m.get("n", 0),
            "return_pct": m.get("return_pct", 0),
            "max_dd_pct": m.get("max_dd_pct", 0),
            "win_rate": m.get("win_rate", 0),
            "cagr_pct": m.get("cagr_pct", 0),
            **extra,
        })
    return pd.DataFrame(rows).sort_values("sharpe", ascending=False)


def main() -> None:
    ap = argparse.ArgumentParser(description="Long VXX puts: sweet-spot roll + VX1 level")
    ap.add_argument("--start", default="2016-01-01")
    ap.add_argument("--end", default="2026-12-31")
    ap.add_argument("--capital", type=float, default=START_CAPITAL)
    ap.add_argument("--risk-budget", type=float, default=3000.0)
    ap.add_argument("--ma-window", type=int, default=60)
    ap.add_argument("--min-roll", type=float, default=SWEET_MIN_ROLL)
    ap.add_argument("--max-roll", type=float, default=SWEET_MAX_ROLL)
    ap.add_argument("--min-elev", type=float, default=SWEET_MIN_ELEV)
    ap.add_argument("--max-elev", type=float, default=SWEET_MAX_ELEV)
    ap.add_argument("--max-vxx-ma", type=float, default=SWEET_MAX_VXX_VS_MA)
    ap.add_argument("--no-vxx-ma-cap", action="store_true")
    ap.add_argument("--sweet-spot-sweep", action="store_true", help="Sweep moneyness × DTE × hold")
    ap.add_argument("--preload-chains", action="store_true")
    ap.add_argument("--out-csv", type=Path, default=LOGS / "vxx_sweet_spot_put_option_sweep.csv")
    ap.add_argument("--out-trades", type=Path, default=LOGS / "vxx_sweet_spot_put_best_trades.jsonl")
    args = ap.parse_args()

    if not CONTANGO_PATH.is_file():
        print(f"Missing {CONTANGO_PATH}; run download_cboe_vix_futures.py", file=sys.stderr)
        sys.exit(1)

    entry = EntryParams(
        min_roll=args.min_roll,
        max_roll=args.max_roll,
        min_elev_vs_ma=args.min_elev,
        max_elev_vs_ma=args.max_elev,
        max_vxx_vs_ma=None if args.no_vxx_ma_cap else args.max_vxx_ma,
        ma_window=args.ma_window,
    )

    ct = enrich_vix_panel(_load_contango(), args.start, args.end, args.ma_window)
    dates = [
        pd.Timestamp(d).normalize()
        for d in sorted(ct.index)
        if args.start <= str(d.date()) <= args.end
    ]
    print(f"Window: {dates[0].date()} → {dates[-1].date()}  ({len(dates)} days)", flush=True)
    print(
        f"Entry: roll∈[{entry.min_roll:.0%},{entry.max_roll or 999:.0%}]  "
        f"vx1_elev∈[{entry.min_elev_vs_ma:.0%},{entry.max_elev_vs_ma:.0%}]  "
        f"max_vxx_vs_ma={entry.max_vxx_vs_ma}",
        flush=True,
    )

    # Count signal days
    mask = pd.Series(False, index=ct.index)
    for d in dates:
        if d in ct.index and entry_signal(ct.loc[d], entry):
            mask.loc[d] = True
    print(f"Sweet-spot signal days: {int(mask.sum())} / {len(dates)}", flush=True)

    if args.preload_chains or args.sweet_spot_sweep:
        from RenTech.strategy_stack.backtest_vxx_vx1_vx3_strategies import _CHAIN_CACHE

        _CHAIN_CACHE.clear()
        print("Pre-loading VXX option chains …", flush=True)
        for j, d in enumerate(dates):
            _load_chain_cached(d)
            if j and j % 400 == 0:
                print(f"  {j}/{len(dates)}", flush=True)

    if args.sweet_spot_sweep:
        print("\nOption sweep (moneyness × expiration × hold) …", flush=True)
        df = run_option_sweep(ct, dates, entry=entry, capital=args.capital, risk_budget=args.risk_budget)
        args.out_csv.parent.mkdir(parents=True, exist_ok=True)
        df.to_csv(args.out_csv, index=False)

        print("\n" + "=" * 100)
        print("TOP 15 BY SHARPE (sweet-spot entries, Theta options)")
        print("=" * 100)
        show = [
            "moneyness_label", "moneyness", "dte_min", "dte_max", "hold_days",
            "n", "sharpe", "return_pct", "max_dd_pct", "win_rate",
            "avg_pnl_usd", "avg_vxx_spot_ret_pct", "median_dte_entry",
        ]
        print(df[show].head(15).to_string(index=False))

        print("\n" + "=" * 100)
        print("TOP 10 BY RETURN (n≥8)")
        print("=" * 100)
        sub = df[df["n"] >= 8].sort_values("return_pct", ascending=False)
        print(sub[show].head(10).to_string(index=False))

        if not df.empty and df.iloc[0]["n"] > 0:
            best = df.iloc[0]
            tr = run_long_put_roll_carry(
                ct,
                dates,
                entry=entry,
                moneyness=float(best["moneyness"]),
                dte_min=int(best["dte_min"]),
                dte_max=int(best["dte_max"]),
                hold_days=int(best["hold_days"]),
                risk_budget_usd=args.risk_budget,
            )
            args.out_trades.parent.mkdir(parents=True, exist_ok=True)
            with args.out_trades.open("w") as f:
                for t in tr:
                    f.write(json.dumps(asdict(t)) + "\n")
            print(f"\nBest config trades → {args.out_trades}")
        print(f"\nFull sweep → {args.out_csv}")
        return

    # Single run: recommended default from forward-return work
    tr = run_long_put_roll_carry(
        ct,
        dates,
        entry=entry,
        moneyness=0.05,
        dte_min=21,
        dte_max=45,
        hold_days=15,
        risk_budget_usd=args.risk_budget,
    )
    m = daily_equity_metrics(tr, dates, args.capital)
    print(f"\nDefault sweet-spot 5% ITM, 21-45 DTE, 15d hold:")
    print(f"  Trades={m['n']}  Sharpe={m['sharpe']:.2f}  Return={m['return_pct']:+.1f}%  "
          f"MaxDD={m['max_dd_pct']:.1f}%  WinRate={m['win_rate']:.1%}")


if __name__ == "__main__":
    main()
