#!/usr/bin/env python3
"""
Systematic exploration of VXX option structures that exploit structural decay
during VIX contango.  Tests multiple strategies side-by-side on the same dates
so results are directly comparable.

Strategies tested
-----------------
1. **bear_call_credit**  — Sell near-ATM call, buy OTM call.  Collect credit,
   keep it when VXX falls or stays flat.  Theta works FOR you.
2. **deep_itm_put**      — Buy high-delta put (~0.80).  Moves almost 1:1 with
   VXX downside but costs less than shorting shares.
3. **ratio_put_1x2**     — Sell 1 ATM put, buy 2 OTM puts.  Small debit or
   credit; accelerating payoff on large VXX drops.
4. **short_call_naked**  — Sell a single OTM call.  Maximum theta capture but
   undefined risk (capped at a practical stop).
5. **put_debit_spread**  — (baseline) Buy ATM put, sell OTM put.

All use the same contango filter, entry frequency, and chain-derived VXX spot.

Example::

    python RenTech/strategy_stack/explore_vxx_decay_strategies.py \\
        --start 2018-06-01 --end 2025-12-31
"""

from __future__ import annotations

import json
import math
import sys
from dataclasses import dataclass, asdict
from pathlib import Path
from typing import Any

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))
DATA_DIR = _REPO / "RenTech" / "data"

from RenTech.strategy_stack.overlay_contract_sizing import resolve_overlay_contracts

THETA_DIR = DATA_DIR / "theta_chunks"
CONTANGO_PATH = DATA_DIR / "vix_futures_cboe.parquet"

SLIPPAGE = 0.005
MULT = 100  # option multiplier


def _load_contango() -> pd.DataFrame:
    ct = pd.read_parquet(CONTANGO_PATH)
    ct.index = pd.to_datetime(ct.index)
    return ct


def _session_date(qt: pd.Series) -> pd.Series:
    qd = pd.to_datetime(qt, utc=False)
    if qd.dt.tz is None:
        qd = qd.dt.tz_localize("America/New_York", ambiguous="NaT", nonexistent="shift_forward")
    else:
        qd = qd.dt.tz_convert("America/New_York")
    return pd.to_datetime(qd.dt.date)


def _load_chain(d: pd.Timestamp) -> pd.DataFrame:
    y, m = d.year, d.month
    path = THETA_DIR / f"vxx_1545_{y:04d}_{m:02d}.parquet"
    if not path.is_file():
        return pd.DataFrame()
    df = pd.read_parquet(path)
    if df.empty:
        return df
    sess = _session_date(df["quote_datetime"])
    df = df.loc[sess == pd.Timestamp(d.date())].copy()
    if df.empty:
        return df
    strike = pd.to_numeric(df["strike"], errors="coerce")
    if float(strike.max(skipna=True)) < 150:
        strike = strike * 10.0
    df["strike"] = strike
    df["mid"] = 0.5 * (pd.to_numeric(df["bid"], errors="coerce") +
                        pd.to_numeric(df["ask"], errors="coerce"))
    df["right_code"] = df["right"].astype(str).str.upper().str.strip().str[0]
    df["expiration_dt"] = pd.to_datetime(df["expiration"]).dt.normalize()
    qt0 = pd.Timestamp(df["quote_datetime"].iloc[0])
    sess_ts = qt0.tz_localize(None).normalize() if qt0.tzinfo is None else qt0.tz_convert(None).normalize()
    exp_naive = df["expiration_dt"].dt.tz_localize(None)
    df["dte"] = (exp_naive - sess_ts).dt.days
    return df


def _spot_from_chain(chain: pd.DataFrame) -> float | None:
    if chain.empty:
        return None
    near = chain[(chain["dte"] >= 7) & (chain["dte"] <= 60)]
    if near.empty:
        return None
    min_exp = near.sort_values("dte")["expiration_dt"].iloc[0]
    atm = near[near["expiration_dt"] == min_exp]
    calls = atm[atm["right_code"] == "C"][["strike", "mid"]]
    puts = atm[atm["right_code"] == "P"][["strike", "mid"]]
    if calls.empty or puts.empty:
        return None
    m = calls.merge(puts, on="strike", suffixes=("_c", "_p"))
    if m.empty:
        return None
    m["s"] = m["strike"] + m["mid_c"] - m["mid_p"]
    m["gap"] = (m["mid_c"] - m["mid_p"]).abs()
    spot = float(m.sort_values("gap").iloc[0]["s"])
    return spot if (math.isfinite(spot) and spot > 0) else None


def _nearest_strike(chain: pd.DataFrame, target: float, right: str, exp: pd.Timestamp) -> pd.Series | None:
    sub = chain[(chain["right_code"] == right) & (chain["expiration_dt"] == exp)]
    if sub.empty:
        return None
    idx = (sub["strike"] - target).abs().idxmin()
    row = sub.loc[idx]
    mid = float(row["mid"])
    if not (math.isfinite(mid) and mid > 0):
        return None
    return row


def _pick_expiry(chain: pd.DataFrame, dte_min: int, dte_max: int) -> pd.Timestamp | None:
    eligible = chain[(chain["dte"] >= dte_min) & (chain["dte"] <= dte_max)]
    if eligible.empty:
        return None
    return eligible.sort_values("dte")["expiration_dt"].iloc[0]


# ---------------------------------------------------------------------------
@dataclass
class Trade:
    strategy: str
    entry_date: str
    exit_date: str
    exit_reason: str
    vxx_entry: float
    vxx_exit: float
    entry_credit_or_debit: float  # positive = net credit received
    exit_value: float
    pnl_total: float
    contango_ratio: float
    vix3m_vix: float
    # Authoritative at entry for this row's qty (max loss or total debit); portfolio merge uses sum().
    broker_risk_usd: float = 0.0
    # Option contracts per standard leg (1x1 vertical: same qty short and long).
    contracts: int = 1
    broker_risk_per_contract_usd: float = 0.0
    underlying: str = "VXX"
    expiration: str = ""
    # Leg strikes for live reconciliation (None if not applicable).
    short_strike: float | None = None
    long_strike: float | None = None
    put_strike: float | None = None
    call_strike: float | None = None
    # When sizing with ``--broker-risk-pct-of-portfolio``: NAV before entry and pct applied.
    nav_at_entry_usd: float = 0.0
    risk_pct_of_portfolio: float = 0.0
    # Audit / reconciliation (optional; defaults keep JSONL rows compact).
    dte_at_entry: int = 0
    dte_at_exit: int = 0
    hold_target_days: int = 0
    days_held: int = 0
    max_loss_one_usd: float | None = None
    max_loss_total_usd: float | None = None
    target_broker_risk_usd: float | None = None
    contracts_requested: int | None = None
    entry_legs_json: str = "{}"


def vxx_built_snapshot(built: dict) -> dict[str, Any]:
    """JSON-serializable copy of a strategy ``built`` dict (floats/ints/strings only)."""
    out: dict[str, Any] = {}
    for k, v in built.items():
        if isinstance(v, (str, int, bool)) or v is None:
            out[k] = v
        elif isinstance(v, float):
            out[k] = v if math.isfinite(v) else None
        else:
            try:
                fv = float(v)
            except (TypeError, ValueError):
                continue
            out[k] = fv if math.isfinite(fv) else None
    return out


def broker_risk_usd_from_built(strategy: str, built: dict, entry_val: float) -> float:
    """Max loss (credit spreads) or debit paid (long premium) for one position."""
    _ = strategy
    ml = built.get("max_loss")
    if ml is not None and math.isfinite(float(ml)) and float(ml) > 0:
        return float(ml)
    deb = built.get("debit")
    if deb is not None and math.isfinite(float(deb)) and float(deb) > 0:
        return float(deb)
    ev = float(entry_val)
    if ev < 0:
        return max(abs(ev), 1.0)
    return max(abs(ev), 1.0)


def resolve_vxx_contracts_and_broker_risk(
    *,
    strategy: str,
    built: dict,
    entry_val: float,
    contracts: int | None,
    target_broker_risk_usd: float | None,
) -> tuple[int, float, float]:
    """
    Returns ``(contracts, broker_risk_per_contract_usd, broker_risk_total_usd)`` for one VXX entry.
    """
    per = max(float(broker_risk_usd_from_built(strategy, built, float(entry_val))), 1e-9)
    return resolve_overlay_contracts(
        contracts=contracts,
        target_broker_risk_usd=target_broker_risk_usd,
        per_contract_broker_risk=per,
    )


def leg_strikes_from_vxx_built(built: dict) -> tuple[float | None, float | None, float | None, float | None]:
    """Return ``(short_strike, long_strike, put_strike, call_strike)`` from ``built`` legs (any strategy)."""
    ss = float(built["short_k"]) if built.get("short_k") is not None else None
    ls = float(built["long_k"]) if built.get("long_k") is not None else None
    ps = cs = None
    if built.get("strike") is not None:
        k = float(built["strike"])
        r = str(built.get("right", "")).upper()
        if r == "P":
            ps = k
        elif r == "C":
            cs = k
    if ss is not None and ls is not None:
        return ss, ls, None, None
    if cs is not None:
        return None, None, None, cs
    if ps is not None:
        return None, None, ps, None
    return ss, ls, ps, cs


def _slipped(price: float, direction: str) -> float:
    """Apply slippage: buying = pay more, selling = receive less."""
    if direction == "buy":
        return price * (1 + SLIPPAGE)
    return price * (1 - SLIPPAGE)


# ---------------------------------------------------------------------------
# Strategy builders: each returns (entry_value, legs_info) or None
# entry_value > 0 means net credit received, < 0 means net debit paid
# legs_info is used to mark exit
# ---------------------------------------------------------------------------

def _build_bear_call_credit(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, width_pct: float, short_mny: float = 1.00):
    """Sell OTM call, buy further OTM call."""
    short_row = _nearest_strike(chain, spot * short_mny, "C", exp)
    long_target = spot * (1.0 + width_pct)
    long_row = _nearest_strike(chain, long_target, "C", exp)
    if short_row is None or long_row is None:
        return None
    sk, lk = float(short_row["strike"]), float(long_row["strike"])
    if lk <= sk:
        return None
    credit = (_slipped(float(short_row["mid"]), "sell") -
              _slipped(float(long_row["mid"]), "buy")) * MULT
    if credit <= 0:
        return None
    max_loss = (lk - sk) * MULT - credit
    return {
        "credit": credit,
        "max_loss": max_loss,
        "max_profit": credit,
        "short_k": sk,
        "long_k": lk,
        "right": "C",
    }


def _build_deep_itm_put(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, delta_target: float):
    """Buy a deep ITM put (strike well above spot)."""
    target_k = spot * (1 + delta_target)  # delta_target ~0.15-0.25 above spot
    row = _nearest_strike(chain, target_k, "P", exp)
    if row is None:
        return None
    k = float(row["strike"])
    if k < spot * 1.05:  # must be meaningfully ITM
        return None
    debit = _slipped(float(row["mid"]), "buy") * MULT
    if debit <= 0:
        return None
    intrinsic = max(k - spot, 0) * MULT
    time_value = debit - intrinsic
    return {
        "debit": debit,
        "strike": k,
        "right": "P",
        "time_value_paid": time_value,
    }


def _build_ratio_put_1x2(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, spread_pct: float):
    """Sell 1 ATM put, buy 2 OTM puts."""
    sell_row = _nearest_strike(chain, spot * 1.00, "P", exp)
    buy_target = spot * (1.0 - spread_pct)
    buy_row = _nearest_strike(chain, buy_target, "P", exp)
    if sell_row is None or buy_row is None:
        return None
    sk, bk = float(sell_row["strike"]), float(buy_row["strike"])
    if bk >= sk:
        return None
    credit_from_sell = _slipped(float(sell_row["mid"]), "sell") * MULT
    debit_from_buys = _slipped(float(buy_row["mid"]), "buy") * MULT * 2
    net = credit_from_sell - debit_from_buys  # positive = net credit
    return {
        "net_entry": net,
        "sell_k": sk,
        "buy_k": bk,
        "sell_count": 1,
        "buy_count": 2,
    }


def _build_short_call(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, otm_pct: float):
    """Sell a single OTM call."""
    target_k = spot * (1 + otm_pct)
    row = _nearest_strike(chain, target_k, "C", exp)
    if row is None:
        return None
    k = float(row["strike"])
    credit = _slipped(float(row["mid"]), "sell") * MULT
    if credit <= 0:
        return None
    return {
        "credit": credit,
        "strike": k,
        "right": "C",
    }


def _build_bull_call_spread(chain: pd.DataFrame, spot: float, exp: pd.Timestamp,
                             long_otm_pct: float, width_pct: float):
    """Bull call debit spread: buy near-ATM call, sell further OTM call."""
    long_row = _nearest_strike(chain, spot * (1 + long_otm_pct), "C", exp)
    short_row = _nearest_strike(chain, spot * (1 + long_otm_pct + width_pct), "C", exp)
    if long_row is None or short_row is None:
        return None
    lk, sk = float(long_row["strike"]), float(short_row["strike"])
    if sk <= lk:
        return None
    debit = (_slipped(float(long_row["mid"]), "buy") -
             _slipped(float(short_row["mid"]), "sell")) * MULT
    if debit <= 0:
        return None
    max_profit = (sk - lk) * MULT - debit
    return {
        "debit": debit,
        "long_k": lk,
        "short_k": sk,
        "max_profit": max_profit,
    }


def _exit_bull_call_spread(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    lk, sk = legs["long_k"], legs["short_k"]
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        lr = sub.iloc[(sub["strike"] - lk).abs().argsort()[:1]]
        sr = sub.iloc[(sub["strike"] - sk).abs().argsort()[:1]]
        if len(lr) and len(sr):
            lm, sm = float(lr.iloc[0]["mid"]), float(sr.iloc[0]["mid"])
            if math.isfinite(lm) and math.isfinite(sm):
                credit = (_slipped(lm, "sell") - _slipped(sm, "buy")) * MULT
                return credit - legs["debit"]
    long_intrin = max(spot - lk, 0) * MULT
    short_intrin = max(spot - sk, 0) * MULT
    return (long_intrin - short_intrin) * (1 - SLIPPAGE) - legs["debit"]


def _build_long_call(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, otm_pct: float):
    """Buy a single OTM call (tail hedge on VXX spikes)."""
    target_k = spot * (1 + otm_pct)
    row = _nearest_strike(chain, target_k, "C", exp)
    if row is None:
        return None
    k = float(row["strike"])
    debit = _slipped(float(row["mid"]), "buy") * MULT
    if debit <= 0:
        return None
    return {
        "debit": debit,
        "strike": k,
        "right": "C",
    }


def _exit_long_call(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    k = legs["strike"]
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        row = sub.iloc[(sub["strike"] - k).abs().argsort()[:1]]
        if len(row):
            m = float(row.iloc[0]["mid"])
            if math.isfinite(m) and m > 0:
                return _slipped(m, "sell") * MULT - legs["debit"]
    intrinsic = max(spot - k, 0) * MULT
    return intrinsic * (1 - SLIPPAGE) - legs["debit"]


def _build_deep_put_long_call(chain: pd.DataFrame, spot: float, exp: pd.Timestamp,
                               put_delta_target: float, call_otm_pct: float):
    """Deep ITM put + long OTM call as tail hedge."""
    put_target_k = spot * (1 + put_delta_target)
    put_row = _nearest_strike(chain, put_target_k, "P", exp)
    if put_row is None:
        return None
    pk = float(put_row["strike"])
    if pk < spot * 1.05:
        return None
    put_debit = _slipped(float(put_row["mid"]), "buy") * MULT
    if put_debit <= 0:
        return None

    call_target_k = spot * (1 + call_otm_pct)
    call_row = _nearest_strike(chain, call_target_k, "C", exp)
    if call_row is None:
        return None
    ck = float(call_row["strike"])
    call_debit = _slipped(float(call_row["mid"]), "buy") * MULT
    if call_debit <= 0:
        return None

    total_debit = put_debit + call_debit
    return {
        "debit": total_debit,
        "put_k": pk,
        "call_k": ck,
        "put_debit": put_debit,
        "call_debit": call_debit,
    }


def _exit_deep_put_long_call(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    pk, ck = legs["put_k"], legs["call_k"]
    put_val = 0.0
    call_val = 0.0

    if not chain.empty:
        put_sub = chain[(chain["right_code"] == "P") & (chain["expiration_dt"] == exp)]
        if not put_sub.empty:
            pr = put_sub.iloc[(put_sub["strike"] - pk).abs().argsort()[:1]]
            if len(pr):
                m = float(pr.iloc[0]["mid"])
                if math.isfinite(m) and m > 0:
                    put_val = _slipped(m, "sell") * MULT

        call_sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)]
        if not call_sub.empty:
            cr = call_sub.iloc[(call_sub["strike"] - ck).abs().argsort()[:1]]
            if len(cr):
                m = float(cr.iloc[0]["mid"])
                if math.isfinite(m) and m > 0:
                    call_val = _slipped(m, "sell") * MULT

    if put_val == 0.0:
        put_val = max(pk - spot, 0) * MULT * (1 - SLIPPAGE)
    if call_val == 0.0:
        call_val = max(spot - ck, 0) * MULT * (1 - SLIPPAGE)

    return (put_val + call_val) - legs["debit"]


def _build_put_debit(chain: pd.DataFrame, spot: float, exp: pd.Timestamp, width_pct: float):
    """Buy ATM put, sell OTM put (baseline)."""
    long_row = _nearest_strike(chain, spot * 1.00, "P", exp)
    short_target = spot * (1.0 - width_pct)
    short_row = _nearest_strike(chain, short_target, "P", exp)
    if long_row is None or short_row is None:
        return None
    lk, sk = float(long_row["strike"]), float(short_row["strike"])
    if sk >= lk:
        return None
    debit = (_slipped(float(long_row["mid"]), "buy") -
             _slipped(float(short_row["mid"]), "sell")) * MULT
    if debit <= 0:
        return None
    max_profit = (lk - sk) * MULT - debit
    return {
        "debit": debit,
        "max_profit": max_profit,
        "long_k": lk,
        "short_k": sk,
    }


# ---------------------------------------------------------------------------
# Exit valuation
# ---------------------------------------------------------------------------

def _exit_bear_call(chain: pd.DataFrame, spot: float, legs: dict) -> float:
    sk, lk = legs["short_k"], legs["long_k"]
    sr = _nearest_strike(chain, sk, "C", None) if not chain.empty else None
    lr = _nearest_strike(chain, lk, "C", None) if not chain.empty else None
    if sr is not None and lr is not None:
        # Cost to close: buy back short, sell long
        cost = (_slipped(float(sr["mid"]), "buy") - _slipped(float(lr["mid"]), "sell")) * MULT
        return legs["credit"] - cost
    # Intrinsic fallback
    short_itm = max(spot - sk, 0) * MULT
    long_itm = max(spot - lk, 0) * MULT
    cost = short_itm - long_itm
    return legs["credit"] - cost


def _exit_bear_call_from_chain(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    sk, lk = legs["short_k"], legs["long_k"]
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        sr = sub.iloc[(sub["strike"] - sk).abs().argsort()[:1]]
        lr = sub.iloc[(sub["strike"] - lk).abs().argsort()[:1]]
        if len(sr) and len(lr):
            sm, lm = float(sr.iloc[0]["mid"]), float(lr.iloc[0]["mid"])
            if math.isfinite(sm) and math.isfinite(lm) and sm > 0:
                cost = (_slipped(sm, "buy") - _slipped(lm, "sell")) * MULT
                return legs["credit"] - cost
    short_itm = max(spot - sk, 0) * MULT
    long_itm = max(spot - lk, 0) * MULT
    return legs["credit"] - (short_itm - long_itm)


def _exit_deep_put(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    k = legs["strike"]
    sub = chain[(chain["right_code"] == "P") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        row = sub.iloc[(sub["strike"] - k).abs().argsort()[:1]]
        if len(row):
            m = float(row.iloc[0]["mid"])
            if math.isfinite(m) and m > 0:
                return _slipped(m, "sell") * MULT - legs["debit"]
    intrinsic = max(k - spot, 0) * MULT
    return intrinsic * (1 - SLIPPAGE) - legs["debit"]


def _exit_ratio_put(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    sk, bk = legs["sell_k"], legs["buy_k"]
    sub = chain[(chain["right_code"] == "P") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        sr = sub.iloc[(sub["strike"] - sk).abs().argsort()[:1]]
        br = sub.iloc[(sub["strike"] - bk).abs().argsort()[:1]]
        if len(sr) and len(br):
            sm, bm = float(sr.iloc[0]["mid"]), float(br.iloc[0]["mid"])
            if math.isfinite(sm) and math.isfinite(bm):
                close_sell = _slipped(sm, "buy") * MULT  # buy back the sold put
                close_buys = _slipped(bm, "sell") * MULT * 2  # sell the 2 long puts
                return legs["net_entry"] - close_sell + close_buys
    sell_intrin = max(sk - spot, 0) * MULT
    buy_intrin = max(bk - spot, 0) * MULT * 2
    return legs["net_entry"] - sell_intrin + buy_intrin * (1 - SLIPPAGE)


def _exit_short_call(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    k = legs["strike"]
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        row = sub.iloc[(sub["strike"] - k).abs().argsort()[:1]]
        if len(row):
            m = float(row.iloc[0]["mid"])
            if math.isfinite(m) and m > 0:
                return legs["credit"] - _slipped(m, "buy") * MULT
    intrinsic = max(spot - k, 0) * MULT
    return legs["credit"] - intrinsic


def _exit_put_debit(chain: pd.DataFrame, spot: float, legs: dict, exp: pd.Timestamp) -> float:
    lk, sk = legs["long_k"], legs["short_k"]
    sub = chain[(chain["right_code"] == "P") & (chain["expiration_dt"] == exp)] if not chain.empty else pd.DataFrame()
    if not sub.empty:
        lr = sub.iloc[(sub["strike"] - lk).abs().argsort()[:1]]
        sr = sub.iloc[(sub["strike"] - sk).abs().argsort()[:1]]
        if len(lr) and len(sr):
            lm, sm = float(lr.iloc[0]["mid"]), float(sr.iloc[0]["mid"])
            if math.isfinite(lm) and math.isfinite(sm):
                credit = (_slipped(lm, "sell") - _slipped(sm, "buy")) * MULT
                return credit - legs["debit"]
    l_intrin = max(lk - spot, 0) * MULT
    s_intrin = max(sk - spot, 0) * MULT
    return (l_intrin - s_intrin) * (1 - SLIPPAGE) - legs["debit"]


# ---------------------------------------------------------------------------
# Main exploration loop
# ---------------------------------------------------------------------------

STRATEGIES = {
    "bear_call_5otm_15w": {
        "builder": lambda c, s, e: _build_bear_call_credit(c, s, e, 0.15, 1.05),
        "exiter": _exit_bear_call_from_chain,
        "entry_field": "credit",
        "max_loss_field": "max_loss",
    },
    "bear_call_atm_10w": {
        "builder": lambda c, s, e: _build_bear_call_credit(c, s, e, 0.10, 1.00),
        "exiter": _exit_bear_call_from_chain,
        "entry_field": "credit",
        "max_loss_field": "max_loss",
    },
    "deep_itm_put": {
        "builder": lambda c, s, e: _build_deep_itm_put(c, s, e, 0.20),
        "exiter": _exit_deep_put,
        "entry_field": None,
        "max_loss_field": None,
    },
    "ratio_put_1x2": {
        "builder": lambda c, s, e: _build_ratio_put_1x2(c, s, e, 0.10),
        "exiter": _exit_ratio_put,
        "entry_field": None,
        "max_loss_field": None,
    },
    "short_call_5pct": {
        "builder": lambda c, s, e: _build_short_call(c, s, e, 0.05),
        "exiter": _exit_short_call,
        "entry_field": "credit",
        "max_loss_field": None,
    },
    "short_call_10pct": {
        "builder": lambda c, s, e: _build_short_call(c, s, e, 0.10),
        "exiter": _exit_short_call,
        "entry_field": "credit",
        "max_loss_field": None,
    },
    "put_debit_spread": {
        "builder": lambda c, s, e: _build_put_debit(c, s, e, 0.08),
        "exiter": _exit_put_debit,
        "entry_field": None,
        "max_loss_field": None,
    },
    "dput20_lcall10": {
        "builder": lambda c, s, e: _build_deep_put_long_call(c, s, e, 0.20, 0.10),
        "exiter": _exit_deep_put_long_call,
        "entry_field": None,
        "max_loss_field": None,
    },
    "dput20_lcall05": {
        "builder": lambda c, s, e: _build_deep_put_long_call(c, s, e, 0.20, 0.05),
        "exiter": _exit_deep_put_long_call,
        "entry_field": None,
        "max_loss_field": None,
    },
    "dput15_lcall10": {
        "builder": lambda c, s, e: _build_deep_put_long_call(c, s, e, 0.15, 0.10),
        "exiter": _exit_deep_put_long_call,
        "entry_field": None,
        "max_loss_field": None,
    },
}


def run_exploration(
    start: str = "2018-06-01",
    end: str = "2025-12-31",
    contango_mode: str = "futures",
    contango_threshold: float = 0.03,
    vix3m_threshold: float = 1.05,
    dte_min: int = 21,
    dte_max: int = 45,
    hold_days: int = 20,
    rebalance_every: int = 10,
    stop_loss_mult: float = 2.0,
) -> dict[str, list[Trade]]:
    ct = _load_contango()
    all_dates = sorted(ct.index)
    dates = [d for d in all_dates if start <= str(d.date()) <= end]
    print(f"Dates: {len(dates)}  ({dates[0].date()} → {dates[-1].date()})  contango={contango_mode}", flush=True)

    results: dict[str, list[Trade]] = {name: [] for name in STRATEGIES}
    pending: dict[str, dict | None] = {name: None for name in STRATEGIES}
    days_held: dict[str, int] = {name: 0 for name in STRATEGIES}
    n = len(dates)
    log_every = max(1, n // 15)

    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))

        if step % log_every == 0:
            sums = {nm: sum(t.pnl_total for t in results[nm]) for nm in STRATEGIES}
            counts = {nm: len(results[nm]) for nm in STRATEGIES}
            line = "  ".join(f"{nm[:10]}={counts[nm]}/${sums[nm]:+.0f}" for nm in STRATEGIES)
            print(f"  [{step:>5}/{n}] {d.date()}  {line}", flush=True)

        # --- exit pass ---
        for name, spec in STRATEGIES.items():
            if pending[name] is None:
                continue
            days_held[name] += 1
            p = pending[name]
            dte_left = int((p["expiration"] - d).days)

            should_exit = (days_held[name] >= p["hold_target"]) or (dte_left <= 1)

            # Early stop-loss check for credit strategies
            if not should_exit and spec["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 = spec["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 not should_exit:
                continue

            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 = spec["exiter"](chain, spot_now, p["legs"], p["expiration"])
            pnl = float(pnl_one) * int(p["contracts"])
            reason = "time" if days_held[name] >= p["hold_target"] or dte_left <= 1 else "stop_loss"
            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

            results[name].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=0.0,
                    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[name]),
                    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=None,
                    contracts_requested=int(p["contracts"]),
                    entry_legs_json=json.dumps(vxx_built_snapshot(p["legs"])),
                )
            )
            pending[name] = None
            days_held[name] = 0

        # --- entry pass ---
        if step % rebalance_every != 0:
            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

        for name, spec in STRATEGIES.items():
            if pending[name] is not None:
                continue  # already in a position
            built = spec["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))))
            )
            n_c, per_u, br_tot = resolve_vxx_contracts_and_broker_risk(
                strategy=name,
                built=built,
                entry_val=entry_val_one,
                contracts=None,
                target_broker_risk_usd=None,
            )
            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

            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[name] = {
                "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_tot,
                "broker_risk_per_contract": per_u,
                "contracts": n_c,
                "dte_at_entry": dte_entry,
                "cr": cr_val if math.isfinite(cr_val) else 0.0,
                "v3v": v3v,
            }
            days_held[name] = 0

    return results


def _summarize(trades: list[Trade]) -> dict:
    if not trades:
        return {"n": 0}
    pnls = [t.pnl_total for t in trades]
    wins = sum(1 for p in pnls if p > 0)
    cum = np.cumsum(pnls)
    peak = np.maximum.accumulate(cum)
    dd = cum - peak
    return {
        "n": len(trades),
        "total_pnl": round(sum(pnls), 2),
        "avg_pnl": round(float(np.mean(pnls)), 2),
        "median_pnl": round(float(np.median(pnls)), 2),
        "win_rate": round(wins / len(trades), 4),
        "best": round(max(pnls), 2),
        "worst": round(min(pnls), 2),
        "max_drawdown": round(float(dd.min()), 2),
        "peak_equity": round(float(peak.max()), 2),
        "sharpe_approx": round(float(np.mean(pnls) / np.std(pnls)) * np.sqrt(26) if np.std(pnls) > 0 else 0, 2),
    }


def main():
    import argparse
    ap = argparse.ArgumentParser(description="Explore VXX decay option strategies")
    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("--out-dir", type=Path, default=DATA_DIR / "logs")
    args = ap.parse_args()

    results = run_exploration(
        start=args.start,
        end=args.end,
        contango_mode=args.contango_mode,
        contango_threshold=args.contango_threshold,
        vix3m_threshold=args.vix3m_threshold,
        hold_days=args.hold_days,
        rebalance_every=args.rebalance_every,
        dte_min=args.dte_min,
        dte_max=args.dte_max,
    )

    print("\n" + "=" * 80)
    print("STRATEGY COMPARISON")
    print("=" * 80)

    summaries = {}
    for name in STRATEGIES:
        s = _summarize(results[name])
        summaries[name] = s
        print(f"\n--- {name} ---")
        print(json.dumps(s, indent=2))

        out_path = args.out_dir / f"vxx_{name}.jsonl"
        out_path.parent.mkdir(parents=True, exist_ok=True)
        with out_path.open("w") as f:
            for t in results[name]:
                f.write(json.dumps(asdict(t)) + "\n")

    # Ranked table
    print("\n" + "=" * 90)
    print(f"{'Strategy':<24} {'Trades':>6} {'Total PnL':>10} {'Avg PnL':>8} {'WinRate':>8} {'Sharpe':>7} {'MaxDD':>8} {'Peak':>8}")
    print("-" * 90)
    ranked = sorted(summaries.items(), key=lambda x: x[1].get("total_pnl", 0), reverse=True)
    for name, s in ranked:
        if s["n"] == 0:
            continue
        print(
            f"{name:<24} {s['n']:>6} {s['total_pnl']:>10.0f} {s['avg_pnl']:>8.1f} "
            f"{s['win_rate']:>7.1%} {s.get('sharpe_approx',0):>7.2f} "
            f"{s['max_drawdown']:>8.0f} {s['peak_equity']:>8.0f}"
        )
    print("=" * 90)


if __name__ == "__main__":
    main()
