#!/usr/bin/env python3
"""
Research backtest: **bear call credit spread** — **sell** a lower-strike call and **buy** a
higher-strike call (same expiry). This is **not** a straddle: it is **short vega** on the premium
you collect (you want IV to fall vs what you sold) and typically **net short delta** (bearish SPY),
with **defined max loss** (width minus credit).

Uses the same XGB **pred_vrp_edge** per strike as the vol-mispricing pipeline. For **rich IV** on
what you **sell**, we rank spreads by the **short (lower) call’s** predicted edge: more **negative**
``pred`` ⇒ model expects **forward RV below IV** on that strike ⇒ better candidate to **sell**
premium.

**Legs (per spread):** short 1× call @ K_lo, long 1× call @ K_hi, ``K_lo < K_hi``, both in the same
expiry / DTE band. Optional filters: both strikes **above** spot (classic OTM bearish credit) via
``--min-call-moneyness``.

**Hedge:** ``hedge_shares = delta_target - net_spread_delta`` in SPY shares (same convention as the
straddle backtest); **negative** ``--delta-target`` ⇒ aim for extra short SPY exposure after hedge.

Example::

    python RenTech/strategy_stack/backtest_iv_rich_vol_bear_call_spread.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --theta-dir RenTech/data/theta_chunks \\
      --start 2020-01-01 --end 2024-01-01 \\
      --max-short-leg-pred 0.05 \\
      --delta-target -20
"""

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
from RenTech.core.theta_chunks_loader import 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.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 BearCallSpreadTradeResult:
    entry_date: str
    exit_date: str
    expiration: str
    strike_short: float
    strike_long: float
    pred_short_leg: float
    pred_long_leg: float
    entry_net_credit: float
    exit_net_debit: 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


def _best_bear_call_spread(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    min_call_moneyness: float,
    max_width_strike: float,
    max_short_leg_pred: float,
) -> tuple[pd.Series, pd.Series, float, float] | None:
    """
    Pick (row_short, row_long, pred_short, pred_long) for the best bear call spread.
    Ranks by **most negative** short-leg pred (IV richest to sell), then narrower width.
    """
    if feat.empty or len(preds) != len(feat):
        return None
    base = df.iloc[feat["_row"].astype(int)].reset_index(drop=True)
    base["pred"] = preds
    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()
    r0 = base["right"].astype(str).str.upper().str.strip().str[0]
    base = base[r0 == "C"].copy()
    if len(base) < 2:
        return None

    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()

    best: tuple[float, float, float, pd.Series, pd.Series] | None = None
    # best: (score_tuple for max, pred_short, pred_long, rs, rl)

    for exp, g in base.groupby("_exp"):
        dte = int((pd.Timestamp(exp).normalize() - sn).days)
        if dte < dte_min or dte > dte_max:
            continue
        strikes = sorted({float(x) for x in g["_k"].values if math.isfinite(float(x))})
        if len(strikes) < 2:
            continue
        for i, k_lo in enumerate(strikes):
            if k_lo < spy_px * min_call_moneyness:
                continue
            for k_hi in strikes[i + 1 :]:
                if k_hi <= k_lo:
                    continue
                if k_hi - k_lo > max_width_strike:
                    break
                if k_hi < spy_px * min_call_moneyness:
                    continue
                kk = g["_k"].astype(float).values
                gs = g[np.isclose(kk, k_lo, rtol=0, atol=1e-4)]
                gl = g[np.isclose(kk, k_hi, rtol=0, atol=1e-4)]
                if len(gs) < 1 or len(gl) < 1:
                    continue
                rs, rl = gs.iloc[0], gl.iloc[0]
                ps, pl = float(rs["pred"]), float(rl["pred"])
                if ps > max_short_leg_pred:
                    continue
                # Prefer more negative ps (richer IV on short leg), then tighter width
                width = k_hi - k_lo
                key = (ps, width)
                if best is None or key < best[0]:
                    best = (key, ps, pl, rs, rl)

    if best is None:
        return None
    _, ps, pl, rs, rl = best
    return rs, rl, ps, pl


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    delta_target: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    min_call_moneyness: float,
    max_width_strike: float,
    max_short_leg_pred: float,
    slippage: float,
) -> list[BearCallSpreadTradeResult]:
    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[BearCallSpreadTradeResult] = []
    pending: dict | None = None

    for i, d in enumerate(trading_days):
        d = pd.Timestamp(d).normalize()
        if pending is not None:
            if d < pending["exit_d"]:
                continue
            c_s = pending["c_short"]
            c_l = pending["c_long"]
            spy_s = float(spy_df.loc[d, "close"])
            # Close: buy back short call, sell long call
            px_s = execution_price_per_share(c_s, "buy", slippage)
            px_l = execution_price_per_share(c_l, "sell", slippage)
            if px_s is None or px_l is None:
                pending = None
                continue
            exit_net = (px_s - px_l) * CONTRACT_MULTIPLIER
            entry_cred = float(pending["entry_credit"])
            pnl_opt = entry_cred - exit_net
            spy_e = float(pending["spy_entry"])
            h = float(pending["hedge_shares"])
            pnl_h = h * (spy_s - spy_e)
            trades.append(
                BearCallSpreadTradeResult(
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration=str(pd.Timestamp(pending["expiration"]).date()),
                    strike_short=float(pending["k_lo"]),
                    strike_long=float(pending["k_hi"]),
                    pred_short_leg=float(pending["pred_s"]),
                    pred_long_leg=float(pending["pred_l"]),
                    entry_net_credit=entry_cred,
                    exit_net_debit=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"]),
                )
            )
            pending = None
            continue

        if i % rebalance_every != 0:
            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)
        picked = _best_bear_call_spread(
            df,
            feat,
            pred,
            spy_px,
            dte_min=dte_min,
            dte_max=dte_max,
            min_call_moneyness=min_call_moneyness,
            max_width_strike=max_width_strike,
            max_short_leg_pred=max_short_leg_pred,
        )
        if picked is None:
            continue
        rs, rl, pred_s, pred_l = picked

        r_rate = 0.04
        rs_s = _with_scaled_strike_row(rs, spy_px)
        rl_s = _with_scaled_strike_row(rl, spy_px)
        c_s = _row_to_contract(rs_s, d, spy_px, r_rate)
        c_l = _row_to_contract(rl_s, d, spy_px, r_rate)
        if c_s is None or c_l is None:
            continue

        # Open: sell lower call, buy higher call
        px_open_s = execution_price_per_share(c_s, "sell", slippage)
        px_open_l = execution_price_per_share(c_l, "buy", slippage)
        if px_open_s is None or px_open_l is None:
            continue
        entry_credit = (px_open_s - px_open_l) * CONTRACT_MULTIPLIER

        ds = float(c_s.delta)
        dl = float(c_l.delta)
        # Short 1 lower call: -100*ds; long 1 upper call: +100*dl
        net_delta = CONTRACT_MULTIPLIER * (dl - ds)
        hedge_shares = float(delta_target) - net_delta

        exp = pd.Timestamp(rs["expiration"]).normalize()
        k_lo = float(c_s.strike)
        k_hi = float(c_l.strike)

        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 = {
            "entry_d": d,
            "exit_d": exit_d,
            "c_short": c_s,
            "c_long": c_l,
            "entry_credit": entry_credit,
            "spy_entry": spy_px,
            "hedge_shares": hedge_shares,
            "expiration": exp,
            "k_lo": k_lo,
            "k_hi": k_hi,
            "pred_s": pred_s,
            "pred_l": pred_l,
            "net_delta": net_delta,
            "delta_target": float(delta_target),
        }

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(
        description="Bear call credit spread: sell rich IV / short delta (not a straddle)",
    )
    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="2020-01-01")
    ap.add_argument("--end", type=str, default="2024-01-01")
    ap.add_argument(
        "--delta-target",
        type=float,
        default=-15.0,
        help="Net SPY delta (share equiv) after hedge; default slightly short (negative).",
    )
    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-call-moneyness",
        type=float,
        default=1.0,
        help="Require K/S >= this for both strikes (1.0 = at/above spot; use >1 for OTM calls only).",
    )
    ap.add_argument(
        "--max-width-strike",
        type=float,
        default=25.0,
        help="Max K_hi - K_lo ($/share strike gap).",
    )
    ap.add_argument(
        "--max-short-leg-pred",
        type=float,
        default=0.0,
        help="Only consider spreads whose **sold** call has pred_vrp_edge <= this (e.g. 0 = RV<=IV).",
    )
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument("--out-trades", type=Path, default=None)
    args = ap.parse_args()

    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,
        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_call_moneyness=float(args.min_call_moneyness),
        max_width_strike=float(args.max_width_strike),
        max_short_leg_pred=float(args.max_short_leg_pred),
        slippage=float(args.slippage),
    )

    pnls = [t.pnl_total for t in trades]
    total = float(np.sum(pnls)) if pnls else 0.0
    win = sum(1 for p in pnls if p > 0)
    print(
        json.dumps(
            {
                "n_trades": len(trades),
                "total_pnl": total,
                "avg_pnl": float(np.mean(pnls)) if pnls else 0.0,
                "win_rate": (win / len(pnls)) if pnls else 0.0,
                "structure": "bear_call_credit_spread",
                "delta_target": float(args.delta_target),
            },
            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()
