#!/usr/bin/env python3
"""
Research backtests using the **vol mispricing** XGB model on Theta Parquet:

**calendar** — Same strike, **two expiries**: **long** far-dated call, **short** near-dated call
(classic long call calendar). Ranks pairs by ``pred_far − pred_near`` (term-structure tilt: prefer
better edge on the long leg vs the short).

**risk_reversal** — Same expiry: **long** OTM put + **short** OTM call. Ranks by ``pred_put − pred_call``
(skew: long cheap-vol put vs short rich-vol call).

**Exit** before the **near** expiry for calendars (short leg binds). **SPY hedge** optional via
``--delta-target``.

JSONL rows include ``pnl_total`` and ``exit_date`` for ``scan_iv_overlay_correlation.py``.

Example::

    python RenTech/strategy_stack/backtest_iv_calendar_risk_reversal.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --mode calendar --out-trades RenTech/data/logs/calendar_spread.jsonl

    python RenTech/strategy_stack/backtest_iv_calendar_risk_reversal.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --mode risk_reversal \\
      --put-moneyness-max 0.98 --call-moneyness-min 1.02 \\
      --out-trades RenTech/data/logs/risk_reversal.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 (
    _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 CalendarRRTrade:
    mode: str
    entry_date: str
    exit_date: str
    expiration_near: str
    expiration_far: str
    strike: float
    pred_signal: float
    entry_net_premium: float
    exit_net_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
    # Approx. broker capital: net debit for calendars; put debit + short-call margin bound for RR.
    broker_risk_usd: float = 0.0
    # Option contracts per leg (same qty on each leg of the structure).
    contracts: int = 1
    broker_risk_per_contract_usd: float = 0.0
    underlying: str = "SPY"
    put_strike: float | None = None
    call_strike: float | None = None
    strike_far: float | None = None
    strike_near: float | None = None
    nav_at_entry_usd: float = 0.0
    risk_pct_of_portfolio: float = 0.0


def _best_calendar_call(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    min_dte_gap: int,
    mny_band: float,
    min_pred_diff: float,
) -> tuple[pd.Series, pd.Series, float] | None:
    """
    Long far call, short near call, same strike. Score = pred_far - pred_near (maximize).
    """
    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 == "C"].copy()
    if base.empty:
        return None
    base["_k"] = [_scale_strike_raw(base.iloc[i], spy_px) for i in range(len(base))]
    base["_exp"] = pd.to_datetime(base["expiration"]).dt.normalize()
    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()

    best: tuple[float, pd.Series, pd.Series] | None = None
    for k, g in base.groupby("_k"):
        m = float(k) / spy_px if spy_px > 0 else float("nan")
        if not math.isfinite(m) or abs(math.log(m)) > mny_band:
            continue
        exps = sorted({pd.Timestamp(x).normalize() for x in g["_exp"].values})
        if len(exps) < 2:
            continue
        for i in range(len(exps)):
            for j in range(i + 1, len(exps)):
                exp_near, exp_far = exps[i], exps[j]
                dte_n = int((exp_near - sn).days)
                dte_f = int((exp_far - sn).days)
                if dte_n < dte_min or dte_n > dte_max or dte_f < dte_min or dte_f > dte_max:
                    continue
                if dte_f - dte_n < min_dte_gap:
                    continue
                gn = g[g["_exp"] == exp_near]
                gf = g[g["_exp"] == exp_far]
                if len(gn) < 1 or len(gf) < 1:
                    continue
                rn, rf = gn.iloc[0], gf.iloc[0]
                pfn, pff = float(rn["pred"]), float(rf["pred"])
                score = pff - pfn
                if score < min_pred_diff:
                    continue
                if best is None or score > best[0]:
                    best = (score, rn, rf)
    if best is None:
        return None
    _, rn, rf = best
    return rn, rf, float(rf["pred"]) - float(rn["pred"])


def _best_risk_reversal(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    put_moneyness_max: float,
    call_moneyness_min: float,
    min_pred_spread: float,
) -> tuple[pd.Series, pd.Series, float] | None:
    """
    Long OTM put, short OTM call, same expiry. Score = pred_put - pred_call (maximize).
    """
    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]
    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()

    best: tuple[float, pd.Series, pd.Series] | None = None
    for exp, g in base.groupby(pd.to_datetime(base["expiration"]).dt.normalize()):
        dte = int((pd.Timestamp(exp).normalize() - sn).days)
        if dte < dte_min or dte > dte_max:
            continue
        rg = g["right"].astype(str).str.upper().str.strip().str[0]
        puts = g[rg == "P"]
        calls = g[rg == "C"]
        if puts.empty or calls.empty:
            continue
        for _, rp in puts.iterrows():
            kp = _scale_strike_raw(rp, spy_px)
            mp = kp / spy_px if spy_px > 0 else float("nan")
            if not math.isfinite(mp) or mp >= put_moneyness_max or mp >= 1.0:
                continue
            for _, rc in calls.iterrows():
                kc = _scale_strike_raw(rc, spy_px)
                mc = kc / spy_px if spy_px > 0 else float("nan")
                if not math.isfinite(mc) or mc <= call_moneyness_min:
                    continue
                score = float(rp["pred"]) - float(rc["pred"])
                if score < min_pred_spread:
                    continue
                if best is None or score > best[0]:
                    best = (score, rp, rc)

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


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    mode: str,
    min_vix: float,
    delta_target: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    min_dte_gap: int,
    mny_band: float,
    min_pred_diff: float,
    put_moneyness_max: float,
    call_moneyness_min: float,
    min_pred_spread: 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[CalendarRRTrade]:
    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[CalendarRRTrade] = []
    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"])
            mode_p = pending["mode"]

            nleg = int(pending["contracts"])
            if mode_p == "calendar":
                c_near = pending["c_near"]
                c_far = pending["c_far"]
                px_sell_far = execution_price_per_share(c_far, "sell", slippage)
                px_buy_near = execution_price_per_share(c_near, "buy", slippage)
                if px_sell_far is None or px_buy_near is None:
                    pending = None
                    continue
                exit_net = (px_sell_far - px_buy_near) * CONTRACT_MULTIPLIER * float(nleg)
            else:
                c_put = pending["c_put"]
                c_call = pending["c_call"]
                px_sell_put = execution_price_per_share(c_put, "sell", slippage)
                px_buy_call = execution_price_per_share(c_call, "buy", slippage)
                if px_sell_put is None or px_buy_call is None:
                    pending = None
                    continue
                exit_net = (px_sell_put - px_buy_call) * CONTRACT_MULTIPLIER * float(nleg)

            entry_net = float(pending["entry_net_premium"])
            pnl_opt = exit_net - entry_net
            pnl_h = h * (spy_s - spy_e)
            realized_pnl += float(pnl_opt + pnl_h)
            trades.append(
                CalendarRRTrade(
                    mode=mode_p,
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration_near=str(pending["exp_near"]),
                    expiration_far=str(pending["exp_far"]),
                    strike=float(pending["strike"]),
                    pred_signal=float(pending["pred_signal"]),
                    entry_net_premium=entry_net,
                    exit_net_premium=exit_net,
                    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",
                    put_strike=pending.get("put_strike"),
                    call_strike=pending.get("call_strike"),
                    strike_far=pending.get("strike_far"),
                    strike_near=pending.get("strike_near"),
                    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 mode == "calendar":
            got = _best_calendar_call(
                df,
                feat,
                pred,
                spy_px,
                dte_min=dte_min,
                dte_max=dte_max,
                min_dte_gap=min_dte_gap,
                mny_band=mny_band,
                min_pred_diff=min_pred_diff,
            )
            if got is None:
                continue
            rn, rf, pred_sig = got
            rn_s = _with_scaled_strike_row(rn, spy_px)
            rf_s = _with_scaled_strike_row(rf, spy_px)
            c_near = _row_to_contract(rn_s, d, spy_px, r_rate)
            c_far = _row_to_contract(rf_s, d, spy_px, r_rate)
            if c_near is None or c_far is None:
                continue
            pay_far = execution_price_per_share(c_far, "buy", slippage)
            px_short_near = execution_price_per_share(c_near, "sell", slippage)
            if pay_far is None or px_short_near is None:
                continue
            entry_net = (pay_far - px_short_near) * CONTRACT_MULTIPLIER
            net_delta = CONTRACT_MULTIPLIER * (float(c_far.delta) - float(c_near.delta))
            hedge_shares = float(delta_target) - net_delta
            exp_near = pd.Timestamp(rn["expiration"]).normalize()
            exp_far = pd.Timestamp(rf["expiration"]).normalize()
            k_strike = float(c_far.strike)
            exit_ideal = _trading_day_offset(spy_idx, d, hold_days)
            last_near = _last_session_before_expiry(spy_idx, d, exp_near)
            if last_near is None or exit_ideal is None:
                continue
            exit_d = min(exit_ideal, last_near)
            if entry_net > 0:
                broker_risk_one = float(entry_net)
            else:
                broker_risk_one = abs(float(entry_net)) + 0.20 * spy_px * CONTRACT_MULTIPLIER
            broker_risk_one = max(broker_risk_one, 1.0)
            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=broker_risk_one,
            )
            entry_net = float(entry_net) * float(n_sz)
            net_delta = float(net_delta) * float(n_sz)
            hedge_shares = float(delta_target) - net_delta
            pending = {
                "mode": "calendar",
                "entry_d": d,
                "exit_d": exit_d,
                "entry_net_premium": entry_net,
                "spy_entry": spy_px,
                "hedge_shares": hedge_shares,
                "pred_signal": pred_sig,
                "net_delta": net_delta,
                "delta_target": float(delta_target),
                "c_near": c_near,
                "c_far": c_far,
                "exp_near": exp_near.date(),
                "exp_far": exp_far.date(),
                "strike": k_strike,
                "contracts": int(n_sz),
                "broker_risk_per_contract": float(per_u),
                "broker_risk_total": float(br_t),
                "put_strike": None,
                "call_strike": None,
                "strike_near": float(c_near.strike),
                "strike_far": float(c_far.strike),
                "nav_at_entry": float(nav_at),
                "risk_pct_of_portfolio": float(pct_applied),
            }
        else:
            got = _best_risk_reversal(
                df,
                feat,
                pred,
                spy_px,
                dte_min=dte_min,
                dte_max=dte_max,
                put_moneyness_max=put_moneyness_max,
                call_moneyness_min=call_moneyness_min,
                min_pred_spread=min_pred_spread,
            )
            if got is None:
                continue
            rp, rc, pred_sig = got
            rp_s = _with_scaled_strike_row(rp, spy_px)
            rc_s = _with_scaled_strike_row(rc, spy_px)
            c_put = _row_to_contract(rp_s, d, spy_px, r_rate)
            c_call = _row_to_contract(rc_s, d, spy_px, r_rate)
            if c_put is None or c_call is None:
                continue
            pay_put = execution_price_per_share(c_put, "buy", slippage)
            px_short_call = execution_price_per_share(c_call, "sell", slippage)
            if pay_put is None or px_short_call is None:
                continue
            entry_net = (pay_put - px_short_call) * CONTRACT_MULTIPLIER
            net_delta = CONTRACT_MULTIPLIER * (float(c_put.delta) - float(c_call.delta))
            hedge_shares = float(delta_target) - net_delta
            exp = pd.Timestamp(rp["expiration"]).normalize()
            k_strike = float(c_put.strike)
            exit_ideal = _trading_day_offset(spy_idx, d, hold_days)
            last_b = _last_session_before_expiry(spy_idx, d, exp)
            if last_b is None or exit_ideal is None:
                continue
            exit_d = min(exit_ideal, last_b)
            put_debit_usd = float(pay_put) * CONTRACT_MULTIPLIER
            call_credit_usd = float(px_short_call) * CONTRACT_MULTIPLIER
            k_call = float(c_call.strike)
            otm_amt = max(spy_px - k_call, 0.0) * CONTRACT_MULTIPLIER
            margin_short_call = max(
                0.25 * spy_px * CONTRACT_MULTIPLIER - otm_amt + call_credit_usd,
                0.10 * spy_px * CONTRACT_MULTIPLIER,
            )
            broker_risk_one = max(put_debit_usd + margin_short_call, 1.0)
            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=broker_risk_one,
            )
            entry_net = float(entry_net) * float(n_sz)
            net_delta = float(net_delta) * float(n_sz)
            hedge_shares = float(delta_target) - net_delta
            pending = {
                "mode": "risk_reversal",
                "entry_d": d,
                "exit_d": exit_d,
                "entry_net_premium": entry_net,
                "spy_entry": spy_px,
                "hedge_shares": hedge_shares,
                "pred_signal": pred_sig,
                "net_delta": net_delta,
                "delta_target": float(delta_target),
                "c_put": c_put,
                "c_call": c_call,
                "exp_near": exp.date(),
                "exp_far": exp.date(),
                "strike": k_strike,
                "contracts": int(n_sz),
                "broker_risk_per_contract": float(per_u),
                "broker_risk_total": float(br_t),
                "put_strike": float(c_put.strike),
                "call_strike": float(c_call.strike),
                "strike_near": None,
                "strike_far": None,
                "nav_at_entry": float(nav_at),
                "risk_pct_of_portfolio": float(pct_applied),
            }

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(description="Calendar spread or risk reversal from IV mispricing XGB")
    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("--mode", choices=("calendar", "risk_reversal"), required=True)
    ap.add_argument("--min-vix", type=float, default=0.0, help="Trade only if VIX > this (0 = no gate)")
    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("--min-dte-gap", type=int, default=7, help="Calendar: min days between expiries")
    ap.add_argument("--moneyness-band", type=float, default=0.15, help="Calendar: |ln(K/S)| <= this")
    ap.add_argument("--min-pred-diff", type=float, default=0.0, help="Calendar: min pred_far - pred_near")
    ap.add_argument("--put-moneyness-max", type=float, default=0.98, help="RR: put K/S < this (below spot)")
    ap.add_argument("--call-moneyness-min", type=float, default=1.02, help="RR: call K/S > this (above spot)")
    ap.add_argument("--min-pred-spread", type=float, default=0.0, help="RR: min pred_put - pred_call")
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument("--contracts", type=int, default=None, help="Contracts per leg (default 1 or from target).")
    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 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,
        mode=str(args.mode),
        min_vix=float(args.min_vix),
        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),
        min_dte_gap=int(args.min_dte_gap),
        mny_band=float(args.moneyness_band),
        min_pred_diff=float(args.min_pred_diff),
        put_moneyness_max=float(args.put_moneyness_max),
        call_moneyness_min=float(args.call_moneyness_min),
        min_pred_spread=float(args.min_pred_spread),
        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),
                "mode": args.mode,
                "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:
                row = asdict(t)
                row["exit_date"] = t.exit_date
                row["pnl_total"] = t.pnl_total
                f.write(json.dumps(row) + "\n")
        print(f"Wrote {len(trades)} trades to {p}")


if __name__ == "__main__":
    main()
