#!/usr/bin/env python3
"""
**Stress long-vol / tail sleeve** (research): only trade when **VIX is above a floor** (default **12**),
then choose a structure using the same **XGB** features as ``train_vol_mispricing_xgb.py``.

**Structures**

* ``straddle`` — Long ATM straddle (call + put same K/exp), ranked by **highest** ``pred_call + pred_pred``
  (same idea as ``backtest_iv_mispricing_straddle.py``).
* ``otm_put`` — Long **one** OTM put (``K/S`` in ``[put_moneyness_min, put_moneyness_max]`` below spot),
  ranked by **highest** single-leg ``pred`` (largest predicted RV − IV on that put = preferred long-vol tail).

**Gate:** ``VIX > --min-vix`` (strictly greater than 12 if you pass ``12``). Skips ultra-low-VIX sessions
where long vol is often a slow bleed.

**Output JSONL** matches ``iv_mispricing_complement.py`` expectations: ``exit_date``, ``pnl_total``, etc.

Example::

    python RenTech/strategy_stack/backtest_iv_stress_long_vol.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --theta-dir RenTech/data/theta_chunks \\
      --start 2016-01-04 --end 2026-04-02 \\
      --min-vix 12 \\
      --structure straddle \\
      --out-trades RenTech/data/logs/stress_longvol_straddle.jsonl

    python RenTech/strategy_stack/backtest_iv_stress_long_vol.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --structure otm_put \\
      --put-moneyness-min 0.88 --put-moneyness-max 0.98 \\
      --out-trades RenTech/data/logs/stress_longvol_otm_put.jsonl
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from dataclasses import asdict, dataclass
from pathlib import Path

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

import joblib
import numpy as np
import pandas as pd

from RenTech.core.theta_chunks_loader import _row_to_contract, ThetaChunksLoader
from RenTech.strategy_stack.backtest_iv_mispricing_straddle import (
    _best_straddle,
    _last_session_before_expiry,
    _load_session_df,
    _scale_strike_raw,
    _trading_day_offset,
    _with_scaled_strike_row,
)
from RenTech.strategy_stack.overlay_contract_sizing import (
    nav_pct_target_and_applied,
    resolve_overlay_contracts,
)
from RenTech.strategy_stack.train_vol_mispricing_xgb import FEATURE_COLUMNS, featurize_dataframe
from RenTech.strategy_stack.vrp_backtester import (
    CONTRACT_MULTIPLIER,
    SLIPPAGE_FACTOR,
    execution_price_per_share,
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)


@dataclass
class StressLongVolTrade:
    structure: str
    entry_date: str
    exit_date: str
    expiration: str
    strike: float
    vix_entry: float
    pred_signal: float
    entry_premium: float
    exit_premium: float
    pnl_options: float
    hedge_shares: float
    spy_entry: float
    spy_exit: float
    pnl_hedge: float
    pnl_total: float
    net_delta_entry: float
    delta_target: float
    # Max loss ≈ abs(entry premium) for long premium (total for ``contracts``).
    broker_risk_usd: float = 0.0
    # Option contracts per leg (straddle: same qty on call and put).
    contracts: int = 1
    broker_risk_per_contract_usd: float = 0.0
    underlying: str = "SPY"
    nav_at_entry_usd: float = 0.0
    risk_pct_of_portfolio: float = 0.0


def _best_otm_put(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    put_moneyness_min: float,
    put_moneyness_max: float,
    min_pred: float,
) -> tuple[pd.Series, float] | None:
    """Long OTM put: K < S ⇒ moneyness K/S in (0, 1). Rank by max pred."""
    if feat.empty or len(preds) != len(feat):
        return None
    base = df.iloc[feat["_row"].astype(int)].reset_index(drop=True)
    base["pred"] = preds
    r = base["right"].astype(str).str.upper().str.strip().str[0]
    base = base[r == "P"].copy()
    if base.empty:
        return None

    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()
    best: tuple[float, pd.Series] | None = None

    for i in range(len(base)):
        row = base.iloc[i]
        k = _scale_strike_raw(row, spy_px)
        if not math.isfinite(k) or spy_px <= 0:
            continue
        m = k / spy_px
        if m >= 1.0 or m < put_moneyness_min or m > put_moneyness_max:
            continue
        exp = pd.Timestamp(row["expiration"]).normalize()
        dte = int((exp - sn).days)
        if dte < dte_min or dte > dte_max:
            continue
        pr = float(row["pred"])
        if pr < min_pred:
            continue
        if best is None or pr > best[0]:
            best = (pr, row)

    if best is None:
        return None
    _, rw = best
    return rw, float(rw["pred"])


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    min_vix: float,
    structure: str,
    delta_target: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    mny_band: float,
    min_edge_sum: float,
    put_moneyness_min: float,
    put_moneyness_max: float,
    min_pred_put: float,
    slippage: 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,
) -> list[StressLongVolTrade]:
    model = bundle["model"]
    cols = list(bundle.get("feature_columns", FEATURE_COLUMNS))
    spy_df = normalize_spy_df(spy_df)
    spy_idx = spy_df.index
    trades: list[StressLongVolTrade] = []
    pending: dict | None = None
    r_rate = 0.04
    realized_pnl = 0.0

    for i, d in enumerate(trading_days):
        d = pd.Timestamp(d).normalize()
        vx = float(spy_df.loc[d, "vix_close"])
        if not (math.isfinite(vx) and vx > float(min_vix)):
            if pending is None:
                continue

        if pending is not None:
            if d < pending["exit_d"]:
                continue
            spy_s = float(spy_df.loc[d, "close"])
            spy_e = float(pending["spy_entry"])
            h = float(pending["hedge_shares"])

            nleg = int(pending["contracts"])
            if pending["structure"] == "straddle":
                c_call = pending["call"]
                c_put = pending["put"]
                px_c = execution_price_per_share(c_call, "sell", slippage)
                px_p = execution_price_per_share(c_put, "sell", slippage)
                if px_c is None or px_p is None:
                    pending = None
                    continue
                exit_prem = (px_c + px_p) * CONTRACT_MULTIPLIER * float(nleg)
            else:
                c_put = pending["put"]
                px_p = execution_price_per_share(c_put, "sell", slippage)
                if px_p is None:
                    pending = None
                    continue
                exit_prem = px_p * CONTRACT_MULTIPLIER * float(nleg)

            entry_prem = float(pending["entry_premium"])
            pnl_opt = exit_prem - entry_prem
            pnl_h = h * (spy_s - spy_e)
            realized_pnl += float(pnl_opt + pnl_h)
            trades.append(
                StressLongVolTrade(
                    structure=str(pending["structure"]),
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration=str(pd.Timestamp(pending["expiration"]).date()),
                    strike=float(pending["strike"]),
                    vix_entry=float(pending["vix_entry"]),
                    pred_signal=float(pending["pred_signal"]),
                    entry_premium=entry_prem,
                    exit_premium=exit_prem,
                    pnl_options=pnl_opt,
                    hedge_shares=h,
                    spy_entry=spy_e,
                    spy_exit=spy_s,
                    pnl_hedge=pnl_h,
                    pnl_total=pnl_opt + pnl_h,
                    net_delta_entry=float(pending["net_delta"]),
                    delta_target=float(pending["delta_target"]),
                    broker_risk_usd=float(pending["broker_risk_total"]),
                    contracts=int(pending["contracts"]),
                    broker_risk_per_contract_usd=float(pending["broker_risk_per_contract"]),
                    underlying="SPY",
                    nav_at_entry_usd=float(pending.get("nav_at_entry", 0.0)),
                    risk_pct_of_portfolio=float(pending.get("risk_pct_of_portfolio", 0.0)),
                )
            )
            pending = None
            continue

        if i % rebalance_every != 0:
            continue

        if not (math.isfinite(vx) and vx > float(min_vix)):
            continue

        df = _load_session_df(theta_dir, d)
        if df.empty:
            continue
        spy_px = float(spy_df.loc[d, "close"])
        feat = featurize_dataframe(df, spy_df)
        if feat.empty:
            continue
        X = feat[cols].astype("float64")
        pred = model.predict(X)

        if structure == "straddle":
            picked = _best_straddle(
                df,
                feat,
                pred,
                spy_px,
                dte_min=dte_min,
                dte_max=dte_max,
                mny_band=mny_band,
                min_edge_sum=min_edge_sum,
            )
            if picked is None:
                continue
            rc, rp, pred_sig = picked
            rc_s = _with_scaled_strike_row(rc, spy_px)
            rp_s = _with_scaled_strike_row(rp, spy_px)
            c_call = _row_to_contract(rc_s, d, spy_px, r_rate)
            c_put = _row_to_contract(rp_s, d, spy_px, r_rate)
            if c_call is None or c_put is None:
                continue
            buy_c = execution_price_per_share(c_call, "buy", slippage)
            buy_p = execution_price_per_share(c_put, "buy", slippage)
            if buy_c is None or buy_p is None:
                continue
            ep1 = (buy_c + buy_p) * CONTRACT_MULTIPLIER
            nd1 = CONTRACT_MULTIPLIER * (float(c_call.delta) + float(c_put.delta))
            nav_at, tgt_nav, pct_applied = nav_pct_target_and_applied(
                initial_portfolio_capital=initial_portfolio_capital,
                realized_pnl_to_date=realized_pnl,
                broker_risk_pct_of_portfolio=broker_risk_pct_of_portfolio,
            )
            n_sz, per_u, br_t = resolve_overlay_contracts(
                contracts=contracts if tgt_nav is None else None,
                target_broker_risk_usd=target_broker_risk_usd if tgt_nav is None else tgt_nav,
                per_contract_broker_risk=ep1,
            )
            entry_prem = ep1 * float(n_sz)
            net_delta = nd1 * float(n_sz)
            hedge_shares = float(delta_target) - net_delta
            exp = pd.Timestamp(rc["expiration"]).normalize()
            k_strike = float(c_call.strike)
            pend_leg = {"call": c_call, "put": c_put}
        elif structure == "otm_put":
            got = _best_otm_put(
                df,
                feat,
                pred,
                spy_px,
                dte_min=dte_min,
                dte_max=dte_max,
                put_moneyness_min=put_moneyness_min,
                put_moneyness_max=put_moneyness_max,
                min_pred=min_pred_put,
            )
            if got is None:
                continue
            rw, pred_sig = got
            rw_s = _with_scaled_strike_row(rw, spy_px)
            c_put = _row_to_contract(rw_s, d, spy_px, r_rate)
            if c_put is None:
                continue
            buy_p = execution_price_per_share(c_put, "buy", slippage)
            if buy_p is None:
                continue
            ep1 = buy_p * CONTRACT_MULTIPLIER
            nd1 = CONTRACT_MULTIPLIER * float(c_put.delta)
            nav_at, tgt_nav, pct_applied = nav_pct_target_and_applied(
                initial_portfolio_capital=initial_portfolio_capital,
                realized_pnl_to_date=realized_pnl,
                broker_risk_pct_of_portfolio=broker_risk_pct_of_portfolio,
            )
            n_sz, per_u, br_t = resolve_overlay_contracts(
                contracts=contracts if tgt_nav is None else None,
                target_broker_risk_usd=target_broker_risk_usd if tgt_nav is None else tgt_nav,
                per_contract_broker_risk=ep1,
            )
            entry_prem = ep1 * float(n_sz)
            net_delta = nd1 * float(n_sz)
            hedge_shares = float(delta_target) - net_delta
            exp = pd.Timestamp(rw["expiration"]).normalize()
            k_strike = float(c_put.strike)
            pend_leg = {"put": c_put}
        else:
            raise ValueError(structure)

        exit_ideal = _trading_day_offset(spy_idx, d, hold_days)
        last_before = _last_session_before_expiry(spy_idx, d, exp)
        if last_before is None or exit_ideal is None:
            continue
        exit_d = min(exit_ideal, last_before)

        pending = {
            "structure": structure,
            "entry_d": d,
            "exit_d": exit_d,
            "entry_premium": entry_prem,
            "spy_entry": spy_px,
            "hedge_shares": hedge_shares,
            "expiration": exp,
            "strike": k_strike,
            "pred_signal": pred_sig,
            "net_delta": net_delta,
            "delta_target": float(delta_target),
            "vix_entry": vx,
            "contracts": int(n_sz),
            "broker_risk_per_contract": float(per_u),
            "broker_risk_total": float(br_t),
            "nav_at_entry": float(nav_at),
            "risk_pct_of_portfolio": float(pct_applied),
            **pend_leg,
        }

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(description="VIX-gated long vol / tail sleeve from IV mispricing model")
    ap.add_argument("--artifact", type=Path, required=True)
    ap.add_argument("--theta-dir", type=Path, default=_REPO / "RenTech" / "data" / "theta_chunks")
    ap.add_argument("--start", type=str, default="2016-01-04")
    ap.add_argument("--end", type=str, default="2026-04-02")
    ap.add_argument(
        "--min-vix",
        type=float,
        default=12.0,
        help="Only trade when VIX is strictly greater than this (e.g. 12).",
    )
    ap.add_argument(
        "--structure",
        choices=("straddle", "otm_put"),
        default="straddle",
        help="straddle = long ATM straddle; otm_put = long single OTM put (tail).",
    )
    ap.add_argument("--delta-target", type=float, default=0.0)
    ap.add_argument("--hold-days", type=int, default=5)
    ap.add_argument("--rebalance-every", type=int, default=5)
    ap.add_argument("--dte-min", type=int, default=7)
    ap.add_argument("--dte-max", type=int, default=60)
    ap.add_argument("--moneyness-band", type=float, default=0.12, help="Straddle: |ln(K/S)| band")
    ap.add_argument("--min-edge-sum", type=float, default=0.0, help="Straddle: min pred_call+pred_put")
    ap.add_argument(
        "--put-moneyness-min",
        type=float,
        default=0.88,
        help="OTM put: min K/S (strike below spot)",
    )
    ap.add_argument("--put-moneyness-max", type=float, default=0.98)
    ap.add_argument(
        "--min-pred-put",
        type=float,
        default=0.0,
        help="OTM put: min predicted edge (RV-IV) to enter",
    )
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument(
        "--contracts",
        type=int,
        default=None,
        help="Option contracts per leg (default 1, or set with --target-broker-risk-usd).",
    )
    ap.add_argument(
        "--target-broker-risk-usd",
        type=float,
        default=None,
        metavar="USD",
        help="Fixed broker-risk budget per entry (contracts implied); ignored if --contracts is set.",
    )
    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 + cumulative realized PnL. "
        "Implies contracts; do not combine with --contracts or --target-broker-risk-usd.",
    )
    ap.add_argument("--out-trades", type=Path, default=None)
    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)

    art = args.artifact.expanduser()
    if not art.is_file():
        print(f"ERROR: not a file: {art}", file=sys.stderr)
        sys.exit(1)

    bundle = joblib.load(art)
    theta_dir = args.theta_dir.expanduser()

    spy_df = normalize_spy_df(
        load_spy_vix_from_yfinance(
            (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d"),
            (pd.Timestamp(args.end) + pd.Timedelta(days=30)).strftime("%Y-%m-%d"),
        )
    )
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_df)
    days = trading_days_intersecting_spy(
        ld,
        spy_df.index,
        pd.Timestamp(args.start),
        pd.Timestamp(args.end),
    )
    if len(days) < 20:
        print("ERROR: too few overlapping trading days", file=sys.stderr)
        sys.exit(1)

    trades = run_backtest(
        theta_dir=theta_dir,
        spy_df=spy_df,
        bundle=bundle,
        trading_days=days,
        min_vix=float(args.min_vix),
        structure=str(args.structure),
        delta_target=float(args.delta_target),
        hold_days=int(args.hold_days),
        rebalance_every=int(args.rebalance_every),
        dte_min=int(args.dte_min),
        dte_max=int(args.dte_max),
        mny_band=float(args.moneyness_band),
        min_edge_sum=float(args.min_edge_sum),
        put_moneyness_min=float(args.put_moneyness_min),
        put_moneyness_max=float(args.put_moneyness_max),
        min_pred_put=float(args.min_pred_put),
        slippage=float(args.slippage),
        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),
    )

    pnls = [t.pnl_total for t in trades]
    print(
        json.dumps(
            {
                "n_trades": len(trades),
                "min_vix": float(args.min_vix),
                "structure": args.structure,
                "total_pnl": float(np.sum(pnls)) if pnls else 0.0,
            },
            indent=2,
        )
    )

    if args.out_trades:
        p = args.out_trades.expanduser()
        p.parent.mkdir(parents=True, exist_ok=True)
        with p.open("w", encoding="utf-8") as f:
            for t in trades:
                f.write(json.dumps(asdict(t)) + "\n")
        print(f"Wrote {len(trades)} trades to {p}")


if __name__ == "__main__":
    main()
