#!/usr/bin/env python3
"""
**Simple SPY OTM bear-call credit spread** — sell a call ~3% OTM, buy a higher call (same expiry).

Complements the stock-only equity book with **short-vol / mild-bearish** convexity:
defined-risk premium collection when SPY is in an uptrend (above SMA200) and VIX is not extreme.

**Entry (when flat):** every ``--rebalance-every`` sessions if:
  * SPY close > SMA(200)
  * ``--vix-min`` <= VIX <= ``--vix-max``
  * Pick expiry in ``[dte_min, dte_max]``; short call ≈ spot × ``--short-moneyness``;
    long call = short + ``--width-dollars`` (default $5 wing).

**Exit:** hold ``--hold-days``, take-profit at ``--take-profit-pct`` of max profit,
stop at ``--stop-loss-pct`` of max loss, or day before expiry.

**Data:** Theta SPY 15:45 chunks (``spy_1545_YYYY_MM.parquet``); VIX/SPY from Yahoo for gates.

Writes daily MTM CSV for ``combine_best_ideas_stack.py``.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_spy_bear_call_spread_standard.py \\
      --start 2016-01-04 --end 2026-06-18 --capital 100000
"""

from __future__ import annotations

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

import numpy as np
import pandas as pd

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

from RenTech.core.theta_chunks_loader import ThetaChunksLoader
from RenTech.strategy_stack.backtest_iv_mispricing_straddle import (
    _load_session_df,
    _scale_strike_raw,
)
from RenTech.strategy_stack.vrp_backtester import (
    CONTRACT_MULTIPLIER,
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_THETA_DIR = _REPO / "RenTech" / "data" / "theta_chunks"
DEFAULT_OUT = LOGS / "spy_bear_call_spread_standard"

SLIPPAGE = 0.005
MULT = CONTRACT_MULTIPLIER


@dataclass
class BearCallTrade:
    entry_date: str
    exit_date: str
    exit_reason: str
    expiration: str
    strike_short: float
    strike_long: float
    spy_entry: float
    spy_exit: float
    entry_credit: float
    exit_cost: float
    pnl_total: float
    max_profit: float
    max_loss: float
    contracts: int
    vix_entry: float


def _prep_chain(df: pd.DataFrame, spy_px: float) -> pd.DataFrame:
    if df.empty:
        return df
    out = df.copy()
    strike = pd.to_numeric(out["strike"], errors="coerce")
    out["strike"] = strike.apply(lambda k: k * 10.0 if math.isfinite(k) and k < 150 else k)
    out["mid"] = 0.5 * (
        pd.to_numeric(out["bid"], errors="coerce") + pd.to_numeric(out["ask"], errors="coerce")
    )
    out["right_code"] = out["right"].astype(str).str.upper().str.strip().str[0]
    out["expiration_dt"] = pd.to_datetime(out["expiration"]).dt.normalize()
    qt0 = pd.Timestamp(out["quote_datetime"].iloc[0])
    sess_ts = qt0.tz_localize(None).normalize() if qt0.tzinfo is None else qt0.tz_convert(None).normalize()
    exp_naive = out["expiration_dt"].dt.tz_localize(None)
    out["dte"] = (exp_naive - sess_ts).dt.days
    out["moneyness"] = out["strike"] / float(spy_px) if spy_px > 0 else np.nan
    return out


def _pick_expiry(chain: pd.DataFrame, dte_min: int, dte_max: int) -> pd.Timestamp | None:
    sub = chain[(chain["dte"] >= dte_min) & (chain["dte"] <= dte_max)]
    if sub.empty:
        return None
    return pd.Timestamp(sub.sort_values("dte")["expiration_dt"].iloc[0])


def _nearest_call(chain: pd.DataFrame, target_strike: float, exp: pd.Timestamp) -> pd.Series | None:
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == exp)]
    if sub.empty:
        return None
    idx = (sub["strike"] - target_strike).abs().idxmin()
    row = sub.loc[idx]
    mid = float(row["mid"])
    return row if math.isfinite(mid) and mid > 0 else None


def _mark_spread(chain: pd.DataFrame, pending: dict) -> float | None:
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == pending["exp"])]
    if sub.empty:
        return None
    sr = sub.iloc[(sub["strike"] - pending["sk"]).abs().argsort()[:1]]
    lr = sub.iloc[(sub["strike"] - pending["lk"]).abs().argsort()[:1]]
    if not len(sr) or not len(lr):
        return None
    sm, lm = float(sr.iloc[0]["mid"]), float(lr.iloc[0]["mid"])
    if not (math.isfinite(sm) and math.isfinite(lm)):
        return None
    cost = (sm * (1 + SLIPPAGE) - lm * (1 - SLIPPAGE)) * MULT
    return pending["credit"] - cost


def _close_spread(chain: pd.DataFrame, spot: float, pending: dict) -> tuple[float, float]:
    sub = chain[(chain["right_code"] == "C") & (chain["expiration_dt"] == pending["exp"])]
    if not sub.empty:
        sr = sub.iloc[(sub["strike"] - pending["sk"]).abs().argsort()[:1]]
        lr = sub.iloc[(sub["strike"] - pending["lk"]).abs().argsort()[:1]]
        if len(sr) and len(lr):
            sm, lm = float(sr.iloc[0]["mid"]), float(lr.iloc[0]["mid"])
            if math.isfinite(sm) and math.isfinite(lm):
                cost = (sm * (1 + SLIPPAGE) - lm * (1 - SLIPPAGE)) * MULT
                return pending["credit"] - cost, cost
    s_itm = max(spot - pending["sk"], 0) * MULT
    l_itm = max(spot - pending["lk"], 0) * MULT
    cost = s_itm - l_itm
    return pending["credit"] - cost, cost


def run_spy_bear_call_mtm(
    dates: list[pd.Timestamp],
    spy_df: pd.DataFrame,
    vix: pd.Series,
    *,
    theta_dir: Path,
    capital: float,
    short_moneyness: float,
    width_dollars: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    vix_min: float,
    vix_max: float,
    take_profit_pct: float,
    stop_loss_pct: float,
    contracts: int,
) -> tuple[list[BearCallTrade], pd.DataFrame]:
    trades: list[BearCallTrade] = []
    pending: dict | None = None
    days_held = 0
    last_entry_idx = -10_000
    cum_realized = 0.0
    prev_mtm = 0.0
    rows: list[dict] = []

    spy_df = normalize_spy_df(spy_df)
    sma200 = spy_df["sma_200"]

    for i, d in enumerate(dates):
        d = pd.Timestamp(d).normalize()
        if d not in spy_df.index:
            continue
        spy_px = float(spy_df.loc[d, "close"])
        sma_v = float(sma200.loc[d]) if d in sma200.index and pd.notna(sma200.loc[d]) else float("nan")
        vx = float(vix.loc[d]) if d in vix.index and pd.notna(vix.loc[d]) else float("nan")

        chain_raw = _load_session_df(theta_dir, d)
        chain = _prep_chain(chain_raw, spy_px) if not chain_raw.empty else pd.DataFrame()

        mtm = 0.0
        margin = 0.0
        daily_pnl = 0.0
        exit_reason = ""

        if pending is not None:
            days_held += 1
            mark = _mark_spread(chain, pending) if not chain.empty else None
            mtm = float(mark) if mark is not None else prev_mtm
            margin = float(pending["max_loss"]) * int(pending["contracts"])
            dte_left = int((pending["exp"] - d).days)
            tp_hit = take_profit_pct > 0 and mtm >= pending["max_profit"] * take_profit_pct
            sl_hit = stop_loss_pct > 0 and mtm <= -pending["max_loss"] * stop_loss_pct
            time_hit = days_held >= pending["ht"] or dte_left <= 1
            if tp_hit or sl_hit or time_hit:
                pnl_per, exit_cost = _close_spread(chain, spy_px, pending)
                reason = "take_profit" if tp_hit else "stop_loss" if sl_hit else "time_or_expiry"
                trade_pnl = pnl_per * int(pending["contracts"])
                cum_realized += trade_pnl
                daily_pnl = trade_pnl - prev_mtm
                trades.append(
                    BearCallTrade(
                        entry_date=str(pending["entry_date"]),
                        exit_date=str(d.date()),
                        exit_reason=reason,
                        expiration=str(pending["exp"].date()),
                        strike_short=float(pending["sk"]),
                        strike_long=float(pending["lk"]),
                        spy_entry=float(pending["spot"]),
                        spy_exit=spy_px,
                        entry_credit=float(pending["credit"]),
                        exit_cost=float(exit_cost),
                        pnl_total=trade_pnl,
                        max_profit=float(pending["max_profit"]) * int(pending["contracts"]),
                        max_loss=float(pending["max_loss"]) * int(pending["contracts"]),
                        contracts=int(pending["contracts"]),
                        vix_entry=float(pending["vix"]),
                    )
                )
                pending = None
                prev_mtm = 0.0
                mtm = 0.0
                margin = 0.0
                exit_reason = reason
            else:
                daily_pnl = mtm - prev_mtm
                prev_mtm = mtm

        elif (i - last_entry_idx) >= rebalance_every:
            gate = (
                math.isfinite(sma_v)
                and spy_px > sma_v
                and math.isfinite(vx)
                and vix_min <= vx <= vix_max
                and not chain.empty
            )
            if gate:
                exp = _pick_expiry(chain, dte_min, dte_max)
                if exp is not None:
                    sk_tgt = spy_px * float(short_moneyness)
                    lk_tgt = sk_tgt + float(width_dollars)
                    short_row = _nearest_call(chain, sk_tgt, exp)
                    long_row = _nearest_call(chain, lk_tgt, exp)
                    if short_row is not None and long_row is not None:
                        sk = float(_scale_strike_raw(short_row, spy_px))
                        lk = float(_scale_strike_raw(long_row, spy_px))
                        if lk > sk:
                            credit = (
                                float(short_row["mid"]) * (1 - SLIPPAGE)
                                - float(long_row["mid"]) * (1 + SLIPPAGE)
                            ) * MULT
                            if credit > 0:
                                max_loss = (lk - sk) * MULT - credit
                                max_profit = credit
                                days_to_exp = sum(1 for dd in dates if d < dd <= exp) - 1
                                ht = min(hold_days, max(days_to_exp, 1))
                                pending = {
                                    "entry_date": d.date(),
                                    "exp": exp,
                                    "sk": sk,
                                    "lk": lk,
                                    "credit": credit,
                                    "max_loss": max_loss,
                                    "max_profit": max_profit,
                                    "spot": spy_px,
                                    "ht": ht,
                                    "vix": vx,
                                    "contracts": int(contracts),
                                }
                                last_entry_idx = i
                                days_held = 0
                                prev_mtm = 0.0
                                margin = max_loss * int(contracts)

        eq = float(capital) + cum_realized + mtm
        rows.append(
            {
                "date": d.strftime("%Y-%m-%d"),
                "daily_ret": daily_pnl / eq if eq > 0 else 0.0,
                "daily_pnl_usd": daily_pnl,
                "equity_usd": eq,
                "mtm_open_usd": mtm,
                "margin_usd": margin,
                "in_position": pending is not None,
                "exit_reason": exit_reason,
            }
        )

    daily = pd.DataFrame(rows)
    if not daily.empty:
        daily["daily_ret"] = daily["equity_usd"].astype(float).pct_change().fillna(0.0)
        daily.loc[daily.index[0], "daily_ret"] = daily["daily_pnl_usd"].iloc[0] / float(capital)
    return trades, daily


def _metrics(daily: pd.DataFrame, *, capital: float) -> dict:
    r = daily["daily_ret"].astype(float)
    eq = daily["equity_usd"].astype(float)
    n = len(r)
    years = n / 252.0
    end = float(eq.iloc[-1])
    tot = end / capital - 1.0
    cagr = (end / capital) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * math.sqrt(252.0)) if sd > 1e-12 else 0.0
    spy_r = load_spy_vix_from_yfinance()["close"].pct_change()
    aligned = pd.DataFrame({"s": r, "spy": spy_r.reindex(pd.to_datetime(daily["date"])).values}).dropna()
    rho = float(aligned.corr().iloc[0, 1]) if len(aligned) > 2 else float("nan")
    return {
        "capital_usd": capital,
        "ending_equity_usd": end,
        "total_return_pct": tot * 100.0,
        "cagr_pct": cagr * 100.0,
        "sharpe": sharpe,
        "max_drawdown_pct": dd * 100.0,
        "n_days": n,
        "corr_spy": rho,
        "pct_days_in_position": float(daily["in_position"].mean()) * 100.0,
    }


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="2026-06-18")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--theta-dir", type=Path, default=DEFAULT_THETA_DIR)
    ap.add_argument("--short-moneyness", type=float, default=1.03, help="Short call strike / spot (default 3%% OTM)")
    ap.add_argument("--width-dollars", type=float, default=5.0, help="Long call strike minus short (default $5)")
    ap.add_argument("--hold-days", type=int, default=21)
    ap.add_argument("--rebalance-every", type=int, default=21)
    ap.add_argument("--dte-min", type=int, default=25)
    ap.add_argument("--dte-max", type=int, default=40)
    ap.add_argument("--vix-min", type=float, default=12.0)
    ap.add_argument("--vix-max", type=float, default=28.0)
    ap.add_argument("--take-profit-pct", type=float, default=0.50)
    ap.add_argument("--stop-loss-pct", type=float, default=0.80)
    ap.add_argument("--contracts", type=int, default=1)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    yf_start = (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    yf_end = (pd.Timestamp(args.end) + pd.Timedelta(days=14)).strftime("%Y-%m-%d")
    spy_df = normalize_spy_df(load_spy_vix_from_yfinance(yf_start, yf_end))
    spy_df.index = pd.to_datetime(spy_df.index).tz_localize(None)
    vix = spy_df["vix_close"].astype(float)
    theta_dir = args.theta_dir.expanduser().resolve()
    ld = ThetaChunksLoader(theta_dir, spy_df=spy_df)
    dates = trading_days_intersecting_spy(
        ld,
        spy_df.index,
        pd.Timestamp(args.start),
        pd.Timestamp(args.end),
    )
    if len(dates) < 50:
        raise SystemExit(f"Too few sessions ({len(dates)})")

    print(
        f"SPY bear-call credit  {args.start} → {args.end}  "
        f"short={args.short_moneyness:.2f}x  width=${args.width_dollars:.0f}  "
        f"hold={args.hold_days}d  VIX [{args.vix_min},{args.vix_max}]",
        flush=True,
    )

    trades, daily = run_spy_bear_call_mtm(
        dates,
        spy_df,
        vix,
        theta_dir=theta_dir,
        capital=float(args.capital),
        short_moneyness=float(args.short_moneyness),
        width_dollars=float(args.width_dollars),
        hold_days=int(args.hold_days),
        rebalance_every=int(args.rebalance_every),
        dte_min=int(args.dte_min),
        dte_max=int(args.dte_max),
        vix_min=float(args.vix_min),
        vix_max=float(args.vix_max),
        take_profit_pct=float(args.take_profit_pct),
        stop_loss_pct=float(args.stop_loss_pct),
        contracts=int(args.contracts),
    )

    m = _metrics(daily, capital=float(args.capital))
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    trades_path = Path(f"{prefix}_trades.csv")
    meta_path = Path(f"{prefix}_meta.json")
    metrics_path = Path(f"{prefix}_metrics.txt")

    daily.to_csv(daily_path, index=False)
    if trades:
        pd.DataFrame([asdict(t) for t in trades]).to_csv(trades_path, index=False)
    meta = {
        "strategy": "spy_bear_call_spread",
        "token": "spy_bear_call",
        **m,
        "n_trades": len(trades),
        "params": {
            "short_moneyness": args.short_moneyness,
            "width_dollars": args.width_dollars,
            "hold_days": args.hold_days,
            "rebalance_every": args.rebalance_every,
            "dte_min": args.dte_min,
            "dte_max": args.dte_max,
            "vix_min": args.vix_min,
            "vix_max": args.vix_max,
            "contracts": args.contracts,
        },
        "daily_csv": str(daily_path),
        "trades_csv": str(trades_path) if trades else None,
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")
    metrics_path.write_text(
        "\n".join(
            [
                f"SPY bear-call spread  {args.start} → {args.end}",
                f"Return {m['total_return_pct']:.1f}%  CAGR {m['cagr_pct']:.1f}%  "
                f"Sharpe {m['sharpe']:.2f}  MaxDD {m['max_drawdown_pct']:.1f}%",
                f"Trades {len(trades)}  ρ(SPY) {m['corr_spy']:.2f}  "
                f"Invested {m['pct_days_in_position']:.0f}% of days",
            ]
        )
        + "\n",
        encoding="utf-8",
    )
    print(
        f"  trades={len(trades)}  return={m['total_return_pct']:.1f}%  "
        f"Sharpe={m['sharpe']:.2f}  maxDD={m['max_drawdown_pct']:.1f}%  "
        f"ρ(SPY)={m['corr_spy']:.2f}",
        flush=True,
    )
    print(f"  daily → {daily_path}", flush=True)


if __name__ == "__main__":
    main()
