#!/usr/bin/env python3
"""
**Long SPY put debit spread overlay** — systematic tail hedge (VIX-spike proxy).

Buys a debit put spread on SPY every ``--rebalance-every`` trading days:

* **Long leg** — OTM put near ``--long-delta`` (default −0.25, ≈ 5% OTM).
* **Short leg** — deeper OTM put near ``--short-delta`` (default −0.07, ≈ 12% OTM).
* **Same expiry** — target ``--dte-target`` (default 30 DTE), ±``--dte-band``.

The spread profits when SPY drops (VIX spike). Max loss = net debit paid; max gain =
spread width − net debit.  No XGBoost signal required — this is a pure structural hedge.

Sizing: fixed ``--contracts`` OR ``--broker-risk-pct-of-portfolio`` × NAV
(broker risk = net debit paid per spread × contracts × 100).

JSONL output is compatible with ``portfolio_vrp_plus_vxx.py`` / ``portfolio_merge_json``.

Example::

    python RenTech/strategy_stack/backtest_long_put_spread_overlay.py \\
      --theta-dir RenTech/data/theta_chunks \\
      --contracts 1 \\
      --out-trades RenTech/data/logs/engine_long_put_spread.jsonl

    python RenTech/strategy_stack/backtest_long_put_spread_overlay.py \\
      --theta-dir RenTech/data/theta_chunks \\
      --broker-risk-pct-of-portfolio 0.005 \\
      --out-trades RenTech/data/logs/engine_long_put_spread_half_pct.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 pandas as pd

from RenTech.core.theta_chunks_loader import ThetaChunksLoader, _row_to_contract
from RenTech.strategy_stack.backtest_iv_mispricing_straddle import (
    _last_session_before_expiry,
    _load_session_df,
    _scale_strike_raw,
    _trading_day_offset,
)
from RenTech.strategy_stack.overlay_contract_sizing import (
    nav_pct_target_and_applied,
    resolve_overlay_contracts,
)
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,
)

_DEFAULT_THETA_DIR = _REPO / "RenTech" / "data" / "theta_chunks"
_DEFAULT_OUT = _REPO / "RenTech" / "data" / "logs" / "engine_long_put_spread.jsonl"


@dataclass
class PutSpreadTrade:
    entry_date: str
    exit_date: str
    expiration: str
    long_strike: float
    short_strike: float
    vix_entry: float
    spy_entry: float
    spy_exit: float
    entry_net_debit: float      # per share, long premium − short premium (positive = debit)
    exit_net_value: float       # per share, value at exit
    pnl_per_share: float        # exit_net_value − entry_net_debit (positive = profit)
    pnl_total: float            # pnl_per_share × contracts × 100
    broker_risk_usd: float      # total debit paid (max loss)
    broker_risk_per_contract_usd: float
    contracts: int
    nav_at_entry_usd: float = 0.0
    risk_pct_of_portfolio: float = 0.0
    underlying: str = "SPY"


def _pick_put(
    df: pd.DataFrame,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    target_delta: float,
    session_date: pd.Timestamp,
) -> pd.Series | None:
    """
    Find a single put row closest to ``target_delta`` (negative, e.g. −0.25) within DTE window.
    Selects the row minimising ``abs(delta − target_delta)`` in absolute value.
    """
    r = df["right"].astype(str).str.upper().str.strip().str[0]
    puts = df[r == "P"].copy()
    if puts.empty:
        return None

    best_row: pd.Series | None = None
    best_score = float("inf")

    for _, row in puts.iterrows():
        exp = pd.Timestamp(row["expiration"]).normalize()
        dte = int((exp - session_date).days)
        if dte < dte_min or dte > dte_max:
            continue
        k = _scale_strike_raw(row, spy_px)
        if not math.isfinite(k) or k >= spy_px:
            continue
        delta_raw = row.get("delta")
        if delta_raw is None or not math.isfinite(float(delta_raw)):
            continue
        d = float(delta_raw)
        if d >= 0:
            continue
        score = abs(d - target_delta)
        if score < best_score:
            best_score = score
            best_row = row

    return best_row


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    trading_days: list[pd.Timestamp],
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    long_delta: float,
    short_delta: float,
    profit_target_mult: 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,
    min_vix: float = 0.0,
    max_vix: float = 80.0,
) -> list[PutSpreadTrade]:
    """
    Run the systematic long put spread overlay.

    Parameters
    ----------
    rebalance_every:
        Open a new spread every N trading days if no position is open.
    long_delta / short_delta:
        Target deltas (both negative). Long = closer to ATM, short = deeper OTM.
    profit_target_mult:
        Close early when net spread value > profit_target_mult × entry debit (e.g. 2.0 = 100% gain).
    min_vix / max_vix:
        Only open new spreads when VIX is inside [min_vix, max_vix].  Default = no filter.
    """
    trading_days_idx = pd.DatetimeIndex([pd.Timestamp(d).normalize() for d in trading_days])
    trades: list[PutSpreadTrade] = []
    realized_pnl = 0.0
    pending: dict | None = None
    last_entry_idx = -rebalance_every  # so first day qualifies
    r_rate = 0.04

    def _exec(row: pd.Series, action: str, spy_close: float, date: pd.Timestamp) -> float | None:
        """Convert raw parquet row → OptionContract → execution price per share."""
        c = _row_to_contract(row, date, spy_close, r_rate)
        if c is None:
            return None
        return execution_price_per_share(c, action, slippage)

    for i, d in enumerate(trading_days):
        d = pd.Timestamp(d).normalize()
        if d not in spy_df.index:
            continue
        spy_row = spy_df.loc[d]
        spy_px = float(spy_row["close"])
        vx = float(spy_row["vix_close"])
        if not (math.isfinite(spy_px) and math.isfinite(vx)):
            continue

        # --- Close open position ---
        if pending is not None:
            exit_now = False
            exit_reason = ""
            if d >= pd.Timestamp(pending["exit_d"]):
                exit_now = True
                exit_reason = "expiry"
            else:
                # Check early profit target
                sess = _load_session_df(theta_dir, d)
                if sess is not None and not sess.empty:
                    lr = _pick_put(
                        sess, spy_px,
                        dte_min=0, dte_max=pending["dte_entry"] + 10,
                        target_delta=pending["long_delta_entry"],
                        session_date=d,
                    )
                    sr = _pick_put(
                        sess, spy_px,
                        dte_min=0, dte_max=pending["dte_entry"] + 10,
                        target_delta=pending["short_delta_entry"],
                        session_date=d,
                    )
                    if lr is not None and sr is not None:
                        lv = _exec(lr, "buy", spy_px, d)
                        sv = _exec(sr, "sell", spy_px, d)
                        if lv is not None and sv is not None:
                            net_val = float(lv) - float(sv)
                            if net_val >= profit_target_mult * pending["entry_debit"]:
                                exit_now = True
                                exit_reason = "profit_target"
                                pending["_exit_net_val"] = net_val

            if exit_now:
                if "_exit_net_val" not in pending:
                    # Exit at expiry — use last available session prices
                    exit_d_final = _last_session_before_expiry(
                        trading_days_idx,
                        pending["entry_d"],
                        pd.Timestamp(pending["expiration"]),
                    )
                    if exit_d_final is None:
                        exit_d_final = d
                    sess2 = _load_session_df(theta_dir, exit_d_final)
                    net_val = 0.0
                    if sess2 is not None and not sess2.empty:
                        lr = _pick_put(
                            sess2, float(spy_df.loc[exit_d_final, "close"]) if exit_d_final in spy_df.index else spy_px,
                            dte_min=0, dte_max=pending["dte_entry"] + 10,
                            target_delta=pending["long_delta_entry"],
                            session_date=exit_d_final,
                        )
                        sr = _pick_put(
                            sess2, float(spy_df.loc[exit_d_final, "close"]) if exit_d_final in spy_df.index else spy_px,
                            dte_min=0, dte_max=pending["dte_entry"] + 10,
                            target_delta=pending["short_delta_entry"],
                            session_date=exit_d_final,
                        )
                        exit_spy = float(spy_df.loc[exit_d_final, "close"]) if exit_d_final in spy_df.index else spy_px
                        lv = _exec(lr, "buy", exit_spy, exit_d_final) if lr is not None else None
                        sv = _exec(sr, "sell", exit_spy, exit_d_final) if sr is not None else None
                        if lv is not None and sv is not None:
                            net_val = float(lv) - float(sv)
                    pending["_exit_net_val"] = net_val

                spy_exit_px = float(spy_df.loc[d, "close"]) if d in spy_df.index else spy_px
                pnl_per_share = pending["_exit_net_val"] - pending["entry_debit"]
                n = pending["contracts"]
                pnl_total = pnl_per_share * n * CONTRACT_MULTIPLIER
                realized_pnl += pnl_total

                trades.append(
                    PutSpreadTrade(
                        entry_date=str(pending["entry_d"].date()),
                        exit_date=str(d.date()),
                        expiration=str(pd.Timestamp(pending["expiration"]).date()),
                        long_strike=pending["long_strike"],
                        short_strike=pending["short_strike"],
                        vix_entry=pending["vix_entry"],
                        spy_entry=pending["spy_entry"],
                        spy_exit=spy_exit_px,
                        entry_net_debit=pending["entry_debit"],
                        exit_net_value=pending["_exit_net_val"],
                        pnl_per_share=pnl_per_share,
                        pnl_total=pnl_total,
                        broker_risk_usd=pending["broker_risk_usd"],
                        broker_risk_per_contract_usd=pending["broker_risk_per_contract_usd"],
                        contracts=n,
                        nav_at_entry_usd=pending.get("nav_at_entry_usd", 0.0),
                        risk_pct_of_portfolio=pending.get("risk_pct_of_portfolio", 0.0),
                    )
                )
                pending = None

        # --- Open new position ---
        if pending is None and (i - last_entry_idx) >= rebalance_every:
            if not (min_vix <= vx <= max_vix):
                continue
            sess = _load_session_df(theta_dir, d)
            if sess is None or sess.empty:
                continue

            session_date = pd.Timestamp(d).normalize()
            lr = _pick_put(sess, spy_px, dte_min=dte_min, dte_max=dte_max,
                           target_delta=long_delta, session_date=session_date)
            sr = _pick_put(sess, spy_px, dte_min=dte_min, dte_max=dte_max,
                           target_delta=short_delta, session_date=session_date)
            if lr is None or sr is None:
                continue

            exp_l = pd.Timestamp(lr["expiration"]).normalize()
            exp_s = pd.Timestamp(sr["expiration"]).normalize()
            if exp_l != exp_s:
                # Require same expiry — prefer long-leg expiry, re-pick short for same expiry
                same_exp_sess = sess[pd.to_datetime(sess["expiration"]).dt.normalize() == exp_l]
                if same_exp_sess.empty:
                    continue
                sr = _pick_put(same_exp_sess, spy_px, dte_min=0, dte_max=dte_max + 10,
                               target_delta=short_delta, session_date=session_date)
                if sr is None:
                    continue

            entry_l = _exec(lr, "buy", spy_px, session_date)
            entry_s = _exec(sr, "sell", spy_px, session_date)
            if entry_l is None or entry_s is None:
                continue
            entry_debit = float(entry_l) - float(entry_s)
            if entry_debit <= 0:
                # Skip: received credit instead of paying debit (unusual, skip to avoid free spread)
                continue

            lk = _scale_strike_raw(lr, spy_px)
            sk = _scale_strike_raw(sr, spy_px)
            spread_width = (lk - sk) * 1.0  # both in $ per share; long_strike > short_strike
            per_contract_risk = entry_debit * CONTRACT_MULTIPLIER

            # Sizing
            nav = initial_portfolio_capital + realized_pnl
            n_broker_target = None
            if broker_risk_pct_of_portfolio is not None:
                _, n_broker_target = nav_pct_target_and_applied(
                    nav_usd=nav,
                    pct_target=broker_risk_pct_of_portfolio,
                )
            n, per_c_risk, total_risk = resolve_overlay_contracts(
                contracts=contracts,
                target_broker_risk_usd=n_broker_target,
                per_contract_broker_risk=per_contract_risk,
            )

            exp_final = pd.Timestamp(lr["expiration"]).normalize()
            dte_entry = int((exp_final - session_date).days)
            exit_target_d = _trading_day_offset(trading_days_idx, d, dte_entry)

            pending = {
                "entry_d": d,
                "exit_d": exit_target_d,
                "expiration": exp_final,
                "long_strike": lk,
                "short_strike": sk,
                "long_delta_entry": long_delta,
                "short_delta_entry": short_delta,
                "dte_entry": dte_entry,
                "vix_entry": vx,
                "spy_entry": spy_px,
                "entry_debit": entry_debit,
                "contracts": n,
                "broker_risk_usd": total_risk,
                "broker_risk_per_contract_usd": per_c_risk,
                "nav_at_entry_usd": nav,
                "risk_pct_of_portfolio": total_risk / max(nav, 1e-6),
            }
            last_entry_idx = i

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(
        description="Systematic long SPY put debit spread overlay (VIX-spike tail hedge).",
        formatter_class=argparse.ArgumentDefaultsHelpFormatter,
    )
    ap.add_argument("--theta-dir", type=Path, default=_DEFAULT_THETA_DIR)
    ap.add_argument("--start", type=str, default="")
    ap.add_argument("--end", type=str, default="")
    ap.add_argument(
        "--rebalance-every",
        type=int,
        default=15,
        metavar="N",
        help="Open a new spread every N trading days (approx monthly = 15).",
    )
    ap.add_argument(
        "--dte-target",
        type=int,
        default=30,
        help="Target DTE for the spread expiry.",
    )
    ap.add_argument(
        "--dte-band",
        type=int,
        default=10,
        help="Acceptable DTE range = [dte_target - dte_band, dte_target + dte_band].",
    )
    ap.add_argument(
        "--long-delta",
        type=float,
        default=-0.25,
        metavar="D",
        help="Target delta for the long (bought) put leg (negative, e.g. -0.25 ≈ 5%% OTM).",
    )
    ap.add_argument(
        "--short-delta",
        type=float,
        default=-0.07,
        metavar="D",
        help="Target delta for the short (sold) put leg (negative, e.g. -0.07 ≈ 12%% OTM).",
    )
    ap.add_argument(
        "--profit-target-mult",
        type=float,
        default=2.0,
        metavar="X",
        help="Close early when spread value > X × entry debit (default: 2.0 = 100%% gain).",
    )
    ap.add_argument(
        "--min-vix",
        type=float,
        default=0.0,
        help="Only open new spreads when VIX >= this value (default: 0 = always).",
    )
    ap.add_argument(
        "--max-vix",
        type=float,
        default=80.0,
        help="Only open new spreads when VIX <= this value (default: 80 = always).",
    )
    ap.add_argument(
        "--contracts",
        type=int,
        default=None,
        help="Fixed contract count (overrides broker-risk-pct-of-portfolio).",
    )
    ap.add_argument(
        "--broker-risk-pct-of-portfolio",
        type=float,
        default=None,
        metavar="FRAC",
        help="Risk budget as fraction of NAV (e.g. 0.005 = 0.5%% of $100k = $500/trade).",
    )
    ap.add_argument(
        "--portfolio-capital",
        type=float,
        default=100_000.0,
        help="Starting NAV for broker_risk_pct sizing.",
    )
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument(
        "--out-trades",
        type=Path,
        default=_DEFAULT_OUT,
        help="JSONL output file.",
    )
    args = ap.parse_args()

    theta_dir = args.theta_dir.expanduser()
    if not theta_dir.is_dir():
        print(f"ERROR: --theta-dir not found: {theta_dir}", file=sys.stderr)
        sys.exit(1)

    if args.contracts is None and args.broker_risk_pct_of_portfolio is None:
        print("Defaulting to --contracts 1 (no sizing arg supplied).")
        contracts: int | None = 1
    else:
        contracts = args.contracts

    dte_min = max(1, args.dte_target - args.dte_band)
    dte_max = args.dte_target + args.dte_band

    print("Loading SPY / VIX panel …")
    # Load full default history; narrow below to actual theta coverage.
    raw_spy = load_spy_vix_from_yfinance()
    spy_df_norm = normalize_spy_df(raw_spy)

    # Build loader to discover actual theta date range.
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_df_norm)
    all_theta_dates = sorted(ld.iter_chain_dates())
    if not all_theta_dates:
        print("ERROR: no chain dates found in theta_dir", file=sys.stderr)
        sys.exit(1)

    t_start = pd.Timestamp(args.start).normalize() if args.start else pd.Timestamp(all_theta_dates[0]).normalize()
    t_end = pd.Timestamp(args.end).normalize() if args.end else pd.Timestamp(all_theta_dates[-1]).normalize()

    trading_days = trading_days_intersecting_spy(ld, spy_df_norm.index, t_start, t_end)
    if len(trading_days) < 5:
        print("ERROR: too few overlapping trading days", file=sys.stderr)
        sys.exit(1)
    print(f"Trading days: {len(trading_days)}  [{trading_days[0].date()} → {trading_days[-1].date()}]")
    print(
        f"Config: rebalance_every={args.rebalance_every}  DTE={dte_min}–{dte_max}  "
        f"long_delta={args.long_delta}  short_delta={args.short_delta}  "
        f"profit_tgt={args.profit_target_mult}×  VIX=[{args.min_vix},{args.max_vix}]"
    )

    trades = run_backtest(
        theta_dir=theta_dir,
        spy_df=spy_df_norm,
        trading_days=trading_days,
        rebalance_every=args.rebalance_every,
        dte_min=dte_min,
        dte_max=dte_max,
        long_delta=args.long_delta,
        short_delta=args.short_delta,
        profit_target_mult=args.profit_target_mult,
        slippage=args.slippage,
        contracts=contracts,
        broker_risk_pct_of_portfolio=args.broker_risk_pct_of_portfolio,
        initial_portfolio_capital=args.portfolio_capital,
        min_vix=args.min_vix,
        max_vix=args.max_vix,
    )

    out_path = args.out_trades.expanduser()
    out_path.parent.mkdir(parents=True, exist_ok=True)
    with out_path.open("w") as fh:
        for t in trades:
            fh.write(json.dumps(asdict(t)) + "\n")

    if trades:
        total_pnl = sum(t.pnl_total for t in trades)
        wins = sum(1 for t in trades if t.pnl_total > 0)
        print(
            f"\nDone: {len(trades)} trades | total PnL ${total_pnl:,.0f} | "
            f"win rate {wins/len(trades):.0%}"
        )
        print(f"Output → {out_path}")
    else:
        print("No trades generated.")


if __name__ == "__main__":
    main()
