#!/usr/bin/env python3
"""
**Iron condor** (short strangle + wings) backtest using the **vol mispricing** XGB model.

Model output ``pred`` ≈ predicted **RV − IV** (forward edge). **High pred** ⇒ IV looks **cheap**
(good for **long** vol). For a **short-vol** condor we want to **sell** relatively **rich** IV on
the body, i.e. **low** pred on the **short** call and short put. Optional: prefer **higher** pred on
**long** wings (cheaper insurance in model units).

**Ranking score** (``--score-mode``):

* ``short_rich`` — maximize ``-(pred_short_call + pred_short_put)`` (default).
* ``body_minus_wings`` — maximize ``-(pred_sc + pred_sp) + (pred_lc + pred_lp)`` (sell rich body,
  buy cheap wings).

**Structure (per expiry):** sell OTM call @ K_sc, sell OTM put @ K_sp; buy long call @ K_lc > K_sc,
long put @ K_lp < K_sp with minimum **strike** gap ``--wing-min-strike`` (SPY points).

**Exit (each session while open):** Prefer **fresh** chain marks for close; fallback to entry contracts.

* **Take profit:** ``--take-profit-pct`` (default **0.5**) — exit when cost to close
  ``<= credit * (1 - pct)`` (≈ 50% of max profit kept at 0.5). Set **0** to disable.
* **Stop:** ``--stop-loss-debit-mult`` (default **2.0**) — exit when cost to close
  ``>= credit * mult``. Set **0** to disable.
* **Time:** latest of ``hold_days`` and last session before expiry (as before).

**Wings:** ``--wing-min-strike`` default **7** SPY points; override per side with
``--wing-min-strike-call`` / ``--wing-min-strike-put``.

Optional SPY hedge via ``--delta-target``.

JSONL: ``exit_date``, ``pnl_total``, ``exit_reason``, plus metadata for ``scan_iv_overlay_correlation.py``.

Example::

    python RenTech/strategy_stack/backtest_iv_iron_condor_xgb.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --out-trades RenTech/data/logs/iron_condor.jsonl \\
      --log-file RenTech/data/logs/iron_condor_run.log

Tail the log in another terminal::

    tail -f RenTech/data/logs/iron_condor_run.log
"""

from __future__ import annotations

import argparse
from collections import Counter
import json
import math
import sys
import time
from dataclasses import asdict, dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import TextIO

_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.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,
)


def _log_line(fp: TextIO | None, msg: str) -> None:
    if fp is None:
        return
    fp.write(msg + "\n")
    fp.flush()


@dataclass
class IronCondorTrade:
    structure: str
    entry_date: str
    exit_date: str
    expiration: str
    strike_short_call: float
    strike_long_call: float
    strike_short_put: float
    strike_long_put: float
    pred_signal: float
    pred_short_call: float
    pred_short_put: float
    pred_long_call: float
    pred_long_put: 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
    exit_reason: str


def _exit_debit_stale_legs(pending: dict, slippage: float) -> float | None:
    """Close using entry-day contract objects (fallback if chain row missing)."""
    c_sc = pending["c_sc"]
    c_sp = pending["c_sp"]
    c_lc = pending["c_lc"]
    c_lp = pending["c_lp"]
    buy_sc = execution_price_per_share(c_sc, "buy", slippage)
    buy_sp = execution_price_per_share(c_sp, "buy", slippage)
    sell_lc = execution_price_per_share(c_lc, "sell", slippage)
    sell_lp = execution_price_per_share(c_lp, "sell", slippage)
    if None in (buy_sc, buy_sp, sell_lc, sell_lp):
        return None
    return (float(buy_sc) + float(buy_sp) - float(sell_lc) - float(sell_lp)) * CONTRACT_MULTIPLIER


def _iron_condor_exit_debit_fresh_chain(
    theta_dir: Path,
    d: pd.Timestamp,
    spy_px: float,
    exp: pd.Timestamp,
    k_sc: float,
    k_lc: float,
    k_sp: float,
    k_lp: float,
    slippage: float,
    r_rate: float,
) -> float | None:
    """Mark-to-close at session ``d`` using that day's Parquet chain (preferred over stale greeks)."""
    df = _load_session_df(theta_dir, d)
    if df.empty:
        return None
    exp_n = pd.Timestamp(exp).normalize()
    r_sc = _lookup_row(df, spy_px, exp_n, k_sc, "C")
    r_lc = _lookup_row(df, spy_px, exp_n, k_lc, "C")
    r_sp = _lookup_row(df, spy_px, exp_n, k_sp, "P")
    r_lp = _lookup_row(df, spy_px, exp_n, k_lp, "P")
    if r_sc is None or r_lc is None or r_sp is None or r_lp is None:
        return None
    r_sc = _with_scaled_strike_row(r_sc, spy_px)
    r_lc = _with_scaled_strike_row(r_lc, spy_px)
    r_sp = _with_scaled_strike_row(r_sp, spy_px)
    r_lp = _with_scaled_strike_row(r_lp, spy_px)
    c_sc = _row_to_contract(r_sc, d, spy_px, r_rate)
    c_sp = _row_to_contract(r_sp, d, spy_px, r_rate)
    c_lc = _row_to_contract(r_lc, d, spy_px, r_rate)
    c_lp = _row_to_contract(r_lp, d, spy_px, r_rate)
    if None in (c_sc, c_sp, c_lc, c_lp):
        return None
    buy_sc = execution_price_per_share(c_sc, "buy", slippage)
    buy_sp = execution_price_per_share(c_sp, "buy", slippage)
    sell_lc = execution_price_per_share(c_lc, "sell", slippage)
    sell_lp = execution_price_per_share(c_lp, "sell", slippage)
    if None in (buy_sc, buy_sp, sell_lc, sell_lp):
        return None
    return (float(buy_sc) + float(buy_sp) - float(sell_lc) - float(sell_lp)) * CONTRACT_MULTIPLIER


def _lookup_row(
    base: pd.DataFrame,
    spy_px: float,
    exp: pd.Timestamp,
    strike: float,
    right: str,
) -> pd.Series | None:
    """First row matching expiration, strike (scaled), and call/put."""
    rg = right.upper()[0]
    exp = pd.Timestamp(exp).normalize()
    for i in range(len(base)):
        r = base.iloc[i]
        if pd.Timestamp(r["expiration"]).normalize() != exp:
            continue
        if str(r["right"]).upper().strip()[0] != rg:
            continue
        k = _scale_strike_raw(r, spy_px)
        if abs(k - float(strike)) < 0.05:
            return r
    return None


def _best_iron_condor(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_min: int,
    dte_max: int,
    call_short_mny_min: float,
    call_short_mny_max: float,
    put_short_mny_min: float,
    put_short_mny_max: float,
    wing_call_pts: float,
    wing_put_pts: float,
    score_mode: str,
    max_short_pred_sum: float | None,
    min_pred_signal: float,
) -> tuple[pd.Series, pd.Series, pd.Series, pd.Series, float, dict[str, float]] | None:
    """
    Returns (r_sc, r_sp, r_lc, r_lp, pred_signal, pred_parts) 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
    sn = pd.Timestamp(pd.Timestamp(base["quote_datetime"].iloc[0]).date()).normalize()

    best: tuple[float, pd.Series, pd.Series, pd.Series, pd.Series, dict[str, float]] | 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]
        calls = g[rg == "C"]
        puts = g[rg == "P"]
        if calls.empty or puts.empty:
            continue

        call_strikes = sorted(
            {_scale_strike_raw(calls.iloc[i], spy_px) for i in range(len(calls))}
        )
        put_strikes = sorted(
            {_scale_strike_raw(puts.iloc[i], spy_px) for i in range(len(puts))}
        )
        if len(call_strikes) < 2 or len(put_strikes) < 2:
            continue

        for k_sc in call_strikes:
            m_sc = k_sc / spy_px if spy_px > 0 else float("nan")
            if (
                not math.isfinite(m_sc)
                or m_sc < call_short_mny_min
                or m_sc > call_short_mny_max
            ):
                continue
            k_lc = None
            for t in call_strikes:
                if t > k_sc and t - k_sc >= wing_call_pts:
                    k_lc = t
                    break
            if k_lc is None:
                continue

            for k_sp in reversed(put_strikes):
                m_sp = k_sp / spy_px if spy_px > 0 else float("nan")
                if (
                    not math.isfinite(m_sp)
                    or m_sp < put_short_mny_min
                    or m_sp > put_short_mny_max
                ):
                    continue
                k_lp = None
                for t in reversed(put_strikes):
                    if t < k_sp and k_sp - t >= wing_put_pts:
                        k_lp = t
                        break
                if k_lp is None:
                    continue

                r_sc = _lookup_row(base, spy_px, exp, k_sc, "C")
                r_lc = _lookup_row(base, spy_px, exp, k_lc, "C")
                r_sp = _lookup_row(base, spy_px, exp, k_sp, "P")
                r_lp = _lookup_row(base, spy_px, exp, k_lp, "P")
                if r_sc is None or r_lc is None or r_sp is None or r_lp is None:
                    continue

                p_sc = float(r_sc["pred"])
                p_sp = float(r_sp["pred"])
                p_lc = float(r_lc["pred"])
                p_lp = float(r_lp["pred"])
                ssum = p_sc + p_sp
                if max_short_pred_sum is not None and ssum > max_short_pred_sum:
                    continue

                if score_mode == "body_minus_wings":
                    sig = -(p_sc + p_sp) + (p_lc + p_lp)
                else:
                    sig = -(p_sc + p_sp)

                if sig < min_pred_signal:
                    continue

                parts = {
                    "pred_short_call": p_sc,
                    "pred_short_put": p_sp,
                    "pred_long_call": p_lc,
                    "pred_long_put": p_lp,
                }
                if best is None or sig > best[0]:
                    best = (sig, r_sc, r_sp, r_lc, r_lp, parts)

    if best is None:
        return None
    sig, r_sc, r_sp, r_lc, r_lp, parts = best
    return r_sc, r_sp, r_lc, r_lp, float(sig), parts


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    min_vix: float | None,
    max_vix: float | None,
    delta_target: float,
    hold_days: int,
    rebalance_every: int,
    dte_min: int,
    dte_max: int,
    call_short_mny_min: float,
    call_short_mny_max: float,
    put_short_mny_min: float,
    put_short_mny_max: float,
    wing_call_pts: float,
    wing_put_pts: float,
    score_mode: str,
    max_short_pred_sum: float | None,
    min_pred_signal: float,
    slippage: float,
    take_profit_pct: float | None,
    stop_loss_debit_mult: float | None,
    log_fp: TextIO | None = None,
    log_every: int = 50,
) -> list[IronCondorTrade]:
    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[IronCondorTrade] = []
    pending: dict | None = None
    r_rate = 0.04
    n_days = len(trading_days)
    t_run0 = time.perf_counter()
    _log_line(
        log_fp,
        f"RUN_START ts={datetime.now(timezone.utc).isoformat()} n_trading_days={n_days} "
        f"rebalance_every={rebalance_every}",
    )

    for i, d in enumerate(trading_days):
        d = pd.Timestamp(d).normalize()
        vx = float(spy_df.loc[d, "vix_close"])
        vix_ok = math.isfinite(vx)

        if log_fp and log_every > 0 and (i % log_every == 0 or i == n_days - 1):
            el = time.perf_counter() - t_run0
            _log_line(
                log_fp,
                f"heartbeat i={i + 1}/{n_days} day={d.date()} elapsed_s={el:.1f} "
                f"trades_closed={len(trades)} pending={pending is not None}",
            )

        if pending is not None:
            entry_d = pd.Timestamp(pending["entry_d"]).normalize()
            max_exit_d = pd.Timestamp(pending["exit_d"]).normalize()
            if d <= entry_d:
                continue

            spy_s = float(spy_df.loc[d, "close"])
            spy_e = float(pending["spy_entry"])
            h = float(pending["hedge_shares"])
            entry_credit = float(pending["entry_net_credit"])

            exit_debit = _iron_condor_exit_debit_fresh_chain(
                theta_dir,
                d,
                spy_s,
                pending["expiration"],
                float(pending["k_sc"]),
                float(pending["k_lc"]),
                float(pending["k_sp"]),
                float(pending["k_lp"]),
                slippage,
                r_rate,
            )
            if exit_debit is None:
                exit_debit = _exit_debit_stale_legs(pending, slippage)

            if exit_debit is None:
                if d >= max_exit_d:
                    _log_line(
                        log_fp,
                        f"WARN exit_skip day={d.date()} no_price; clearing pending",
                    )
                    pending = None
                continue

            tp_hit = (
                take_profit_pct is not None
                and take_profit_pct > 0.0
                and exit_debit <= entry_credit * (1.0 - take_profit_pct)
            )
            sl_hit = (
                stop_loss_debit_mult is not None
                and stop_loss_debit_mult > 0.0
                and exit_debit >= entry_credit * stop_loss_debit_mult
            )
            time_hit = d >= max_exit_d

            if not (time_hit or tp_hit or sl_hit):
                continue

            if tp_hit:
                exit_reason = "take_profit"
            elif sl_hit:
                exit_reason = "stop_loss"
            else:
                exit_reason = "time"

            pnl_opt = entry_credit - exit_debit
            pnl_h = h * (spy_s - spy_e)

            trades.append(
                IronCondorTrade(
                    structure="iron_condor",
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration=str(pd.Timestamp(pending["expiration"]).date()),
                    strike_short_call=float(pending["k_sc"]),
                    strike_long_call=float(pending["k_lc"]),
                    strike_short_put=float(pending["k_sp"]),
                    strike_long_put=float(pending["k_lp"]),
                    pred_signal=float(pending["pred_signal"]),
                    pred_short_call=float(pending["pred_short_call"]),
                    pred_short_put=float(pending["pred_short_put"]),
                    pred_long_call=float(pending["pred_long_call"]),
                    pred_long_put=float(pending["pred_long_put"]),
                    entry_net_credit=entry_credit,
                    exit_net_debit=exit_debit,
                    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"]),
                    exit_reason=exit_reason,
                )
            )
            _log_line(
                log_fp,
                f"EXIT_TRADE day={d.date()} reason={exit_reason} pnl_total={pnl_opt + pnl_h:.2f} "
                f"debit={exit_debit:.2f} credit={entry_credit:.2f} trades_closed={len(trades)}",
            )
            pending = None
            continue

        if i % rebalance_every != 0:
            continue

        if min_vix is not None and (not vix_ok or vx <= float(min_vix)):
            continue
        if max_vix is not None and (not vix_ok or vx >= float(max_vix)):
            continue

        t0 = time.perf_counter()
        df = _load_session_df(theta_dir, d)
        t1 = time.perf_counter()
        if df.empty:
            _log_line(log_fp, f"rebalance day={d.date()} skip=empty_parquet ms_load={(t1 - t0) * 1000:.1f}")
            continue
        spy_px = float(spy_df.loc[d, "close"])
        feat = featurize_dataframe(df, spy_df)
        t2 = time.perf_counter()
        if feat.empty:
            _log_line(
                log_fp,
                f"rebalance day={d.date()} skip=empty_features n_df={len(df)} "
                f"ms_load={(t1 - t0) * 1000:.1f} ms_feat={(t2 - t1) * 1000:.1f}",
            )
            continue
        X = feat[cols].astype("float64")
        pred = model.predict(X)
        t3 = time.perf_counter()

        got = _best_iron_condor(
            df,
            feat,
            pred,
            spy_px,
            dte_min=dte_min,
            dte_max=dte_max,
            call_short_mny_min=call_short_mny_min,
            call_short_mny_max=call_short_mny_max,
            put_short_mny_min=put_short_mny_min,
            put_short_mny_max=put_short_mny_max,
            wing_call_pts=wing_call_pts,
            wing_put_pts=wing_put_pts,
            score_mode=score_mode,
            max_short_pred_sum=max_short_pred_sum,
            min_pred_signal=min_pred_signal,
        )
        t4 = time.perf_counter()
        _log_line(
            log_fp,
            f"rebalance day={d.date()} n_df={len(df)} n_feat={len(feat)} "
            f"ms_load={(t1 - t0) * 1000:.1f} ms_feat={(t2 - t1) * 1000:.1f} "
            f"ms_pred={(t3 - t2) * 1000:.1f} ms_ic={(t4 - t3) * 1000:.1f} ic_ok={got is not None}",
        )
        if got is None:
            continue
        r_sc, r_sp, r_lc, r_lp, pred_sig, parts = got

        r_sc_s = _with_scaled_strike_row(r_sc, spy_px)
        r_sp_s = _with_scaled_strike_row(r_sp, spy_px)
        r_lc_s = _with_scaled_strike_row(r_lc, spy_px)
        r_lp_s = _with_scaled_strike_row(r_lp, spy_px)

        c_sc = _row_to_contract(r_sc_s, d, spy_px, r_rate)
        c_sp = _row_to_contract(r_sp_s, d, spy_px, r_rate)
        c_lc = _row_to_contract(r_lc_s, d, spy_px, r_rate)
        c_lp = _row_to_contract(r_lp_s, d, spy_px, r_rate)
        if None in (c_sc, c_sp, c_lc, c_lp):
            continue

        sell_sc = execution_price_per_share(c_sc, "sell", slippage)
        sell_sp = execution_price_per_share(c_sp, "sell", slippage)
        buy_lc = execution_price_per_share(c_lc, "buy", slippage)
        buy_lp = execution_price_per_share(c_lp, "buy", slippage)
        if None in (sell_sc, sell_sp, buy_lc, buy_lp):
            continue

        entry_credit = (
            float(sell_sc) + float(sell_sp) - float(buy_lc) - float(buy_lp)
        ) * CONTRACT_MULTIPLIER

        net_delta = CONTRACT_MULTIPLIER * (
            float(c_sc.delta)
            + float(c_sp.delta)
            - float(c_lc.delta)
            - float(c_lp.delta)
        )
        hedge_shares = float(delta_target) - net_delta

        exp = pd.Timestamp(r_sc["expiration"]).normalize()
        k_sc = _scale_strike_raw(r_sc, spy_px)
        k_lc = _scale_strike_raw(r_lc, spy_px)
        k_sp = _scale_strike_raw(r_sp, spy_px)
        k_lp = _scale_strike_raw(r_lp, spy_px)

        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)

        pending = {
            "entry_d": d,
            "exit_d": exit_d,
            "entry_net_credit": entry_credit,
            "spy_entry": spy_px,
            "hedge_shares": hedge_shares,
            "expiration": exp,
            "k_sc": k_sc,
            "k_lc": k_lc,
            "k_sp": k_sp,
            "k_lp": k_lp,
            "pred_signal": pred_sig,
            "pred_short_call": parts["pred_short_call"],
            "pred_short_put": parts["pred_short_put"],
            "pred_long_call": parts["pred_long_call"],
            "pred_long_put": parts["pred_long_put"],
            "net_delta": net_delta,
            "delta_target": float(delta_target),
            "c_sc": c_sc,
            "c_sp": c_sp,
            "c_lc": c_lc,
            "c_lp": c_lp,
        }
        _log_line(
            log_fp,
            f"ENTER_TRADE entry={d.date()} exit={exit_d.date()} pred_signal={pred_sig:.4f} "
            f"k_sc={k_sc:.1f} k_lc={k_lc:.1f} k_sp={k_sp:.1f} k_lp={k_lp:.1f}",
        )

    _log_line(
        log_fp,
        f"RUN_DONE elapsed_s={time.perf_counter() - t_run0:.1f} trades={len(trades)}",
    )
    return trades


def main() -> None:
    ap = argparse.ArgumentParser(description="Iron condor backtest ranked by XGB vol-mispricing preds")
    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(
        "--min-vix",
        type=float,
        default=None,
        help="If set, only trade when VIX > this (e.g. 12). Default: no floor.",
    )
    ap.add_argument(
        "--max-vix",
        type=float,
        default=None,
        help="If set, only trade when VIX < this (e.g. 28). Default: no cap.",
    )
    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=45)
    ap.add_argument(
        "--call-short-mny-min",
        type=float,
        default=1.015,
        help="Short call: min K/S (OTM call)",
    )
    ap.add_argument("--call-short-mny-max", type=float, default=1.045)
    ap.add_argument("--put-short-mny-min", type=float, default=0.92)
    ap.add_argument("--put-short-mny-max", type=float, default=0.985)
    ap.add_argument(
        "--wing-min-strike",
        type=float,
        default=7.0,
        help="Default min strike gap (SPY pts) call wing and put wing; overridden by *_call / *_put.",
    )
    ap.add_argument(
        "--wing-min-strike-call",
        type=float,
        default=None,
        help="Override: min SPY points between short call and long call (default: --wing-min-strike).",
    )
    ap.add_argument(
        "--wing-min-strike-put",
        type=float,
        default=None,
        help="Override: min SPY points between short put and long put (default: --wing-min-strike).",
    )
    ap.add_argument(
        "--take-profit-pct",
        type=float,
        default=0.5,
        help="Exit when cost to close <= credit * (1 - this); 0.5 ≈ 50%% of max profit. Use 0 to disable.",
    )
    ap.add_argument(
        "--stop-loss-debit-mult",
        type=float,
        default=2.0,
        help="Exit when cost to close >= credit * this (e.g. 2.0 = pay 2× credit to close). Use 0 to disable.",
    )
    ap.add_argument(
        "--score-mode",
        choices=("short_rich", "body_minus_wings"),
        default="short_rich",
        help="short_rich: max -(pred_sc+pred_sp); body_minus_wings: add (pred_lc+pred_lp)",
    )
    ap.add_argument(
        "--max-short-pred-sum",
        type=float,
        default=None,
        help="Optional cap: skip if pred_short_call + pred_short_put exceeds this (tighter = richer shorts only)",
    )
    ap.add_argument(
        "--min-pred-signal",
        type=float,
        default=-1e9,
        help="Minimum ranking score to enter (after transform above)",
    )
    ap.add_argument("--slippage", type=float, default=SLIPPAGE_FACTOR)
    ap.add_argument("--out-trades", type=Path, default=None)
    ap.add_argument(
        "--log-file",
        type=Path,
        default=None,
        help="Append human-readable progress (heartbeats, per-rebalance timings, trades).",
    )
    ap.add_argument(
        "--log-every",
        type=int,
        default=50,
        help="Heartbeat every N trading-day iterations (0 = heartbeats off). Default 50.",
    )
    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)

    mss = args.max_short_pred_sum
    if mss is not None:
        mss = float(mss)

    w_base = float(args.wing_min_strike)
    wing_call_pts = float(args.wing_min_strike_call) if args.wing_min_strike_call is not None else w_base
    wing_put_pts = float(args.wing_min_strike_put) if args.wing_min_strike_put is not None else w_base

    tp = float(args.take_profit_pct)
    if tp <= 0.0:
        tp = None
    slm = float(args.stop_loss_debit_mult)
    if slm <= 0.0:
        slm = None

    log_fp: TextIO | None = None
    log_path = args.log_file.expanduser() if args.log_file else None
    if log_path is not None:
        log_path.parent.mkdir(parents=True, exist_ok=True)
        log_fp = log_path.open("a", encoding="utf-8", buffering=1)
        _log_line(
            log_fp,
            f"CLI_START argv={sys.argv[1:]} n_days={len(days)} "
            f"wing_call={wing_call_pts} wing_put={wing_put_pts} "
            f"take_profit_pct={tp} stop_loss_debit_mult={slm}",
        )

    try:
        trades = run_backtest(
            theta_dir=theta_dir,
            spy_df=spy_df,
            bundle=bundle,
            trading_days=days,
            min_vix=float(args.min_vix) if args.min_vix is not None else None,
            max_vix=float(args.max_vix) if args.max_vix is not None else None,
            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),
            call_short_mny_min=float(args.call_short_mny_min),
            call_short_mny_max=float(args.call_short_mny_max),
            put_short_mny_min=float(args.put_short_mny_min),
            put_short_mny_max=float(args.put_short_mny_max),
            wing_call_pts=wing_call_pts,
            wing_put_pts=wing_put_pts,
            score_mode=str(args.score_mode),
            max_short_pred_sum=mss,
            min_pred_signal=float(args.min_pred_signal),
            slippage=float(args.slippage),
            take_profit_pct=tp,
            stop_loss_debit_mult=slm,
            log_fp=log_fp,
            log_every=max(0, int(args.log_every)),
        )
    finally:
        if log_fp is not None:
            log_fp.close()
            print(f"Log appended: {log_path}", flush=True)

    pnls = [t.pnl_total for t in trades]
    ex_ct = Counter(t.exit_reason for t in trades)
    print(
        json.dumps(
            {
                "n_trades": len(trades),
                "structure": "iron_condor",
                "score_mode": args.score_mode,
                "total_pnl": float(np.sum(pnls)) if pnls else 0.0,
                "exit_reason_counts": dict(ex_ct),
                "wing_call_pts": wing_call_pts,
                "wing_put_pts": wing_put_pts,
                "take_profit_pct": tp,
                "stop_loss_debit_mult": slm,
            },
            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()
