#!/usr/bin/env python3
"""
Research backtest: **long ATM straddles** ranked by XGBoost **pred_vrp_edge** (IV mispricing signal),
with a **SPY share hedge** so net option + stock delta matches ``--delta-target`` (default **0** =
delta-neutral in share equivalents).

For each rebalance session:
  * load Theta Parquet (prefers ``*_ivfilled``),
  * score rows with the trained ``joblib`` model,
  * pair **call + put** at the same **expiration** and **strike** (straddle),
  * rank by ``pred_call + pred_put`` (sum of predicted forward RV − IV edge),
  * pick the best straddle in the DTE / moneyness band,
  * enter long 1× call + 1× put at **buy** prices; **hedge_shares = delta_target − net_option_delta**,
    where ``net_option_delta = 100 × (Δ_call + Δ_put)`` per contract.

**Net SPY:** total delta in share terms is ``net_option_delta + hedge_shares`` = ``delta_target``.
So **negative** ``--delta-target`` ⇒ **net short SPY** after hedge; **positive** ⇒ net long. Straddles
often have slightly **negative** net option delta (puts a bit more negative than calls are positive
near ATM), so ``delta_target=0`` usually means **long SPY** hedge — use ``--delta-target-bear`` when
``pred_call + pred_put`` is at/below ``--bear-edge-threshold`` to lean **short SPY** on “IV rich” edges.

Exit after ``--hold-days`` **trading** sessions (clamped to the last session **before** expiry).

**Limitations:** static hedge (no intraday rebalance), Theta 15:45 snapshots only, no commissions,
no margin; for research only.

Example::

    python RenTech/strategy_stack/backtest_iv_mispricing_straddle.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --theta-dir RenTech/data/theta_chunks \\
      --start 2020-01-01 --end 2024-01-01 \\
      --delta-target 0 --hold-days 5 --rebalance-every 5

    # When model says straddle edge is weak / negative (IV rich vs RV), target net short SPY:
    python ... --delta-target 0 --delta-target-bear -35 --bear-edge-threshold 0
"""

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,
    _session_dates_series,
)
from RenTech.core.theta_chunks_loader import ThetaChunksLoader
from RenTech.strategy_stack.train_vol_mispricing_xgb import (
    FEATURE_COLUMNS,
    featurize_dataframe,
    resolve_month_parquet,
)
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 StraddleTradeResult:
    entry_date: str
    exit_date: str
    expiration: str
    strike: float
    pred_edge_sum: 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


def _with_scaled_strike_row(r: pd.Series, spy_px: float) -> pd.Series:
    o = r.copy()
    o["strike"] = _scale_strike_raw(o, spy_px)
    return o


def _scale_strike_raw(r: pd.Series, spy_px: float) -> float:
    k = float(pd.to_numeric(r["strike"], errors="coerce"))
    if not math.isfinite(k):
        return float("nan")
    if k < 150:
        k *= 10.0
    return float(k)


def _load_session_df(theta_dir: Path, session: pd.Timestamp) -> pd.DataFrame:
    path = resolve_month_parquet(theta_dir, session)
    if path is None:
        return pd.DataFrame()
    df = pd.read_parquet(path)
    if df.empty:
        return df
    qd = _session_dates_series(df["quote_datetime"])
    mask = pd.to_datetime(qd, errors="coerce").dt.date == session.date()
    return df.loc[mask].copy()


def _best_straddle(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    mny_band: float,
    min_edge_sum: float,
) -> tuple[pd.Series, pd.Series, float] | None:
    """
    Returns (call_row, put_row, pred_sum) for the best straddle, or None.
    """
    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()
    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()

    best: tuple[float, pd.Series, pd.Series] | None = None
    for (_exp, k), g in base.groupby(["_exp", "_k"]):
        if len(g) < 2:
            continue
        r = g["right"].astype(str).str.upper().str.strip().str[0]
        gc = g[r == "C"]
        gp = g[r == "P"]
        if len(gc) != 1 or len(gp) != 1:
            continue
        rc, rp = gc.iloc[0], gp.iloc[0]
        dte = int((pd.Timestamp(_exp).normalize() - sn).days)
        if dte < dte_min or dte > dte_max:
            continue
        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
        ps = float(rc["pred"]) + float(rp["pred"])
        if ps < min_edge_sum:
            continue
        if best is None or ps > best[0]:
            best = (ps, rc, rp)
    if best is None:
        return None
    _, rc, rp = best
    return rc, rp, float(rc["pred"]) + float(rp["pred"])


def _trading_day_offset(spy_index: pd.DatetimeIndex, d0: pd.Timestamp, off: int) -> pd.Timestamp | None:
    """``d0`` is normalized; ``off`` is number of trading days forward on ``spy_index``."""
    d0 = pd.Timestamp(d0).normalize()
    if d0 not in spy_index:
        return None
    pos = spy_index.get_loc(d0)
    if isinstance(pos, slice):
        pos = pos.start
    j = int(pos) + int(off)
    if j < 0 or j >= len(spy_index):
        return None
    return pd.Timestamp(spy_index[j]).normalize()


def _last_session_before_expiry(
    spy_index: pd.DatetimeIndex,
    entry: pd.Timestamp,
    expiration: pd.Timestamp,
) -> pd.Timestamp | None:
    """Last trading day in ``spy_index`` with date < expiration."""
    exp = pd.Timestamp(expiration).normalize()
    eligible = [d for d in spy_index if d >= entry and d < exp]
    return pd.Timestamp(eligible[-1]).normalize() if eligible else None


def _resolve_delta_target(
    pred_sum: float,
    *,
    delta_target: float,
    delta_target_bear: float | None,
    bear_edge_threshold: float,
) -> float:
    """Use *bear* net SPY target when combined predicted edge is at/below threshold (IV rich)."""
    if delta_target_bear is not None and pred_sum <= bear_edge_threshold:
        return float(delta_target_bear)
    return float(delta_target)


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    delta_target: float,
    delta_target_bear: float | None,
    bear_edge_threshold: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    mny_band: float,
    min_edge_sum: float,
    slippage: float,
) -> list[StraddleTradeResult]:
    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[StraddleTradeResult] = []
    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
            # Close at exit session
            c_call = pending["call"]
            c_put = pending["put"]
            spy_s = float(spy_df.loc[d, "close"])
            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 or not math.isfinite(px_c + px_p):
                pending = None
                continue
            exit_prem = (px_c + px_p) * CONTRACT_MULTIPLIER
            entry_prem = float(pending["entry_premium"])
            pnl_opt = exit_prem - entry_prem
            spy_e = float(pending["spy_entry"])
            spy_x = spy_s
            h = float(pending["hedge_shares"])
            pnl_h = h * (spy_x - spy_e)
            trades.append(
                StraddleTradeResult(
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration=str(pd.Timestamp(pending["expiration"]).date()),
                    strike=float(pending["strike"]),
                    pred_edge_sum=float(pending["pred_sum"]),
                    entry_premium=entry_prem,
                    exit_premium=exit_prem,
                    pnl_options=pnl_opt,
                    hedge_shares=h,
                    spy_entry=spy_e,
                    spy_exit=spy_x,
                    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_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_sum = picked

        r_rate = 0.04
        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
        entry_prem = (buy_c + buy_p) * CONTRACT_MULTIPLIER

        dc = float(c_call.delta)
        dp = float(c_put.delta)
        net_delta = CONTRACT_MULTIPLIER * (dc + dp)
        dt = _resolve_delta_target(
            float(pred_sum),
            delta_target=delta_target,
            delta_target_bear=delta_target_bear,
            bear_edge_threshold=bear_edge_threshold,
        )
        hedge_shares = float(dt) - net_delta

        exp = pd.Timestamp(rc["expiration"]).normalize()
        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:
            continue
        if exit_ideal is None:
            continue
        exit_d = min(exit_ideal, last_before)

        pending = {
            "entry_d": d,
            "exit_d": exit_d,
            "call": c_call,
            "put": c_put,
            "entry_premium": entry_prem,
            "spy_entry": spy_px,
            "hedge_shares": hedge_shares,
            "expiration": exp,
            "strike": float(c_call.strike),
            "pred_sum": pred_sum,
            "net_delta": net_delta,
            "delta_target": float(dt),
        }

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(description="IV mispricing straddle + SPY delta hedge backtest")
    ap.add_argument("--artifact", type=Path, required=True, help="vol_mispricing_xgb.joblib")
    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=0.0,
        help="Net SPY delta in share equivalents after hedge (options+stock). Negative = net short SPY.",
    )
    ap.add_argument(
        "--delta-target-bear",
        type=float,
        default=None,
        help="If set: when pred_call+pred_put <= --bear-edge-threshold, use this net delta instead "
        "(typically negative for net short SPY on 'IV rich' edges).",
    )
    ap.add_argument(
        "--bear-edge-threshold",
        type=float,
        default=0.0,
        help="Bear / short-SPY tilt activates when combined predicted edge <= this (default 0).",
    )
    ap.add_argument("--hold-days", type=int, default=5, help="Trading days to hold (clamped before expiry)")
    ap.add_argument("--rebalance-every", type=int, default=5, help="Try a new entry every N trading days")
    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="Log-moneyness: |ln(K/S)| <= this")
    ap.add_argument("--min-edge-sum", type=float, default=0.0, help="Min pred_call + pred_put to enter")
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument("--out-trades", type=Path, default=None, help="Write trades JSONL")
    ap.add_argument("--capital", type=float, default=100_000.0)
    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),
        delta_target_bear=args.delta_target_bear,
        bear_edge_threshold=float(args.bear_edge_threshold),
        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),
        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)
    n_bear = (
        sum(1 for t in trades if t.pred_edge_sum <= float(args.bear_edge_threshold))
        if args.delta_target_bear is not None
        else 0
    )
    summary = {
        "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,
        "delta_target": float(args.delta_target),
        "delta_target_bear": args.delta_target_bear,
        "bear_edge_threshold": float(args.bear_edge_threshold),
        "n_trades_pred_edge_below_bear_threshold": n_bear,
        "hold_days": int(args.hold_days),
    }
    print(json.dumps(summary, 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()
