#!/usr/bin/env python3
"""
**Flyagonal** (research): combine

1. **Short put calendar** — **long** far-dated put, **short** near-dated put, **same strike** ``K_p``.
   Calendar gap between expiries: ``--calendar-gap-min`` … ``--calendar-gap-max`` calendar days (default 7–14).

2. **Call butterfly** — same **near** expiration as the short put: **long** lower call, **short 2×** middle,
   **long** upper call (equal strike width ``--fly-strike-width``).

Front-week DTE is controlled by ``--dte-near-min`` / ``--dte-near-max`` (default ~7). The far put uses
``--dte-far-min`` / ``--dte-far-max`` so that ``exp_far - exp_near`` is in the gap band.

**XGB ranking** (default): maximize ``(pred_far_put - pred_near_put) + (pred_c_lo + pred_c_hi - 2*pred_c_mid)``.

**No SPY hedge** — ``pnl_total`` = options P&amp;L only (``hedge_shares=0``).

JSONL includes ``exit_date`` and ``pnl_total`` for ``scan_iv_overlay_correlation.py``.

Example::

    python RenTech/strategy_stack/backtest_iv_flyagonal_xgb.py \\
      --artifact RenTech/data/models/vol_mispricing_xgb.joblib \\
      --out-trades RenTech/data/logs/flyagonal.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 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,
)


@dataclass
class FlyagonalTrade:
    structure: str
    entry_date: str
    exit_date: str
    expiration_near: str
    expiration_far: str
    strike_put: float
    strike_call_low: float
    strike_call_mid: float
    strike_call_high: float
    pred_signal: float
    pred_far_put: float
    pred_near_put: float
    pred_c_lo: float
    pred_c_mid: float
    pred_c_hi: float
    entry_net_premium: float
    exit_net_premium: float
    pnl_options: float
    hedge_shares: float
    spy_entry: float
    spy_exit: float
    pnl_hedge: float
    pnl_total: float


def _lookup_row(
    base: pd.DataFrame,
    spy_px: float,
    exp: pd.Timestamp,
    strike: float,
    right: str,
) -> pd.Series | None:
    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_flyagonal(
    df: pd.DataFrame,
    feat: pd.DataFrame,
    preds: np.ndarray,
    spy_px: float,
    *,
    dte_near_min: int,
    dte_near_max: int,
    dte_far_min: int,
    dte_far_max: int,
    calendar_gap_min: int,
    calendar_gap_max: int,
    put_moneyness_min: float,
    put_moneyness_max: float,
    fly_strike_width: float,
    min_pred_signal: float,
) -> tuple[pd.Series, pd.Series, pd.Series, pd.Series, pd.Series, float, dict[str, float]] | None:
    """
    Returns (r_far_put, r_near_put, r_c_lo, r_c_mid, r_c_hi, pred_signal, 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()

    exps = sorted(
        {pd.Timestamp(x).normalize() for x in pd.to_datetime(base["expiration"]).dt.normalize().unique()}
    )
    best: tuple[float, pd.Series, pd.Series, pd.Series, pd.Series, pd.Series, dict[str, float]] | None = None

    for exp_near in exps:
        dte_n = int((exp_near - sn).days)
        if dte_n < dte_near_min or dte_n > dte_near_max:
            continue
        for exp_far in exps:
            if exp_far <= exp_near:
                continue
            gap = int((exp_far - exp_near).days)
            if gap < calendar_gap_min or gap > calendar_gap_max:
                continue
            dte_f = int((exp_far - sn).days)
            if dte_f < dte_far_min or dte_f > dte_far_max:
                continue

            gn = base[pd.to_datetime(base["expiration"]).dt.normalize() == exp_near]
            gf = base[pd.to_datetime(base["expiration"]).dt.normalize() == exp_far]
            rgn = gn["right"].astype(str).str.upper().str.strip().str[0]
            rgf = gf["right"].astype(str).str.upper().str.strip().str[0]
            puts_n = gn[rgn == "P"]
            puts_f = gf[rgf == "P"]
            calls_n = gn[rgn == "C"]
            if puts_n.empty or puts_f.empty or calls_n.empty:
                continue

            put_strikes = sorted(
                {_scale_strike_raw(puts_n.iloc[i], spy_px) for i in range(len(puts_n))}
            )
            call_strikes = sorted(
                {_scale_strike_raw(calls_n.iloc[i], spy_px) for i in range(len(calls_n))}
            )
            if len(call_strikes) < 3:
                continue

            for k_p in put_strikes:
                mp = k_p / spy_px if spy_px > 0 else float("nan")
                if not math.isfinite(mp) or mp < put_moneyness_min or mp > put_moneyness_max:
                    continue
                r_far = _lookup_row(base, spy_px, exp_far, k_p, "P")
                r_near = _lookup_row(base, spy_px, exp_near, k_p, "P")
                if r_far is None or r_near is None:
                    continue

                for i in range(len(call_strikes) - 2):
                    k_lo = call_strikes[i]
                    k_mid = call_strikes[i + 1]
                    k_hi = call_strikes[i + 2]
                    if abs((k_mid - k_lo) - fly_strike_width) > 0.25:
                        continue
                    if abs((k_hi - k_mid) - fly_strike_width) > 0.25:
                        continue

                    r_lo = _lookup_row(base, spy_px, exp_near, k_lo, "C")
                    r_md = _lookup_row(base, spy_px, exp_near, k_mid, "C")
                    r_hi = _lookup_row(base, spy_px, exp_near, k_hi, "C")
                    if r_lo is None or r_md is None or r_hi is None:
                        continue

                    p_far = float(r_far["pred"])
                    p_near = float(r_near["pred"])
                    p_lo = float(r_lo["pred"])
                    p_md = float(r_md["pred"])
                    p_hi = float(r_hi["pred"])
                    sig = (p_far - p_near) + (p_lo + p_hi - 2.0 * p_md)
                    if sig < min_pred_signal:
                        continue

                    parts = {
                        "pred_far_put": p_far,
                        "pred_near_put": p_near,
                        "pred_c_lo": p_lo,
                        "pred_c_mid": p_md,
                        "pred_c_hi": p_hi,
                    }
                    if best is None or sig > best[0]:
                        best = (sig, r_far, r_near, r_lo, r_md, r_hi, parts)

    if best is None:
        return None
    sig, r_far, r_near, r_lo, r_md, r_hi, parts = best
    return r_far, r_near, r_lo, r_md, r_hi, float(sig), parts


def run_backtest(
    *,
    theta_dir: Path,
    spy_df: pd.DataFrame,
    bundle: dict,
    trading_days: list[pd.Timestamp],
    min_vix: float,
    hold_days: int,
    rebalance_every: int,
    dte_near_min: int,
    dte_near_max: int,
    dte_far_min: int,
    dte_far_max: int,
    calendar_gap_min: int,
    calendar_gap_max: int,
    put_moneyness_min: float,
    put_moneyness_max: float,
    fly_strike_width: float,
    min_pred_signal: float,
    slippage: float,
) -> list[FlyagonalTrade]:
    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[FlyagonalTrade] = []
    pending: dict | None = None
    r_rate = 0.04

    for i, d in enumerate(trading_days):
        d = pd.Timestamp(d).normalize()
        vx = float(spy_df.loc[d, "vix_close"])
        if not (math.isfinite(vx) and vx > float(min_vix)):
            if pending is None:
                continue

        if pending is not None:
            if d < pending["exit_d"]:
                continue
            spy_s = float(spy_df.loc[d, "close"])
            spy_e = float(pending["spy_entry"])

            c_fp = pending["c_far_put"]
            c_np = pending["c_near_put"]
            c_lo = pending["c_lo"]
            c_md = pending["c_mid"]
            c_hi = pending["c_hi"]

            sell_fp = execution_price_per_share(c_fp, "sell", slippage)
            buy_np = execution_price_per_share(c_np, "buy", slippage)
            sell_lo = execution_price_per_share(c_lo, "sell", slippage)
            buy_md = execution_price_per_share(c_md, "buy", slippage)
            sell_hi = execution_price_per_share(c_hi, "sell", slippage)
            if None in (sell_fp, buy_np, sell_lo, buy_md, sell_hi):
                pending = None
                continue
            exit_net = (
                float(sell_fp)
                - float(buy_np)
                + float(sell_lo)
                - 2.0 * float(buy_md)
                + float(sell_hi)
            ) * CONTRACT_MULTIPLIER

            entry_net = float(pending["entry_net_premium"])
            pnl_opt = exit_net - entry_net

            trades.append(
                FlyagonalTrade(
                    structure="flyagonal",
                    entry_date=str(pending["entry_d"].date()),
                    exit_date=str(d.date()),
                    expiration_near=str(pending["exp_near"]),
                    expiration_far=str(pending["exp_far"]),
                    strike_put=float(pending["k_p"]),
                    strike_call_low=float(pending["k_cl"]),
                    strike_call_mid=float(pending["k_cm"]),
                    strike_call_high=float(pending["k_ch"]),
                    pred_signal=float(pending["pred_signal"]),
                    pred_far_put=float(pending["pred_far_put"]),
                    pred_near_put=float(pending["pred_near_put"]),
                    pred_c_lo=float(pending["pred_c_lo"]),
                    pred_c_mid=float(pending["pred_c_mid"]),
                    pred_c_hi=float(pending["pred_c_hi"]),
                    entry_net_premium=entry_net,
                    exit_net_premium=exit_net,
                    pnl_options=pnl_opt,
                    hedge_shares=0.0,
                    spy_entry=spy_e,
                    spy_exit=spy_s,
                    pnl_hedge=0.0,
                    pnl_total=pnl_opt,
                )
            )
            pending = None
            continue

        if i % rebalance_every != 0:
            continue
        if not (math.isfinite(vx) and vx > float(min_vix)):
            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)

        got = _best_flyagonal(
            df,
            feat,
            pred,
            spy_px,
            dte_near_min=dte_near_min,
            dte_near_max=dte_near_max,
            dte_far_min=dte_far_min,
            dte_far_max=dte_far_max,
            calendar_gap_min=calendar_gap_min,
            calendar_gap_max=calendar_gap_max,
            put_moneyness_min=put_moneyness_min,
            put_moneyness_max=put_moneyness_max,
            fly_strike_width=fly_strike_width,
            min_pred_signal=min_pred_signal,
        )
        if got is None:
            continue
        r_fp, r_np, r_lo, r_md, r_hi, pred_sig, parts = got

        r_fp_s = _with_scaled_strike_row(r_fp, spy_px)
        r_np_s = _with_scaled_strike_row(r_np, spy_px)
        r_lo_s = _with_scaled_strike_row(r_lo, spy_px)
        r_md_s = _with_scaled_strike_row(r_md, spy_px)
        r_hi_s = _with_scaled_strike_row(r_hi, spy_px)

        c_fp = _row_to_contract(r_fp_s, d, spy_px, r_rate)
        c_np = _row_to_contract(r_np_s, d, spy_px, r_rate)
        c_lo = _row_to_contract(r_lo_s, d, spy_px, r_rate)
        c_md = _row_to_contract(r_md_s, d, spy_px, r_rate)
        c_hi = _row_to_contract(r_hi_s, d, spy_px, r_rate)
        if None in (c_fp, c_np, c_lo, c_md, c_hi):
            continue

        pay_fp = execution_price_per_share(c_fp, "buy", slippage)
        recv_np = execution_price_per_share(c_np, "sell", slippage)
        pay_lo = execution_price_per_share(c_lo, "buy", slippage)
        recv_md = execution_price_per_share(c_md, "sell", slippage)
        pay_hi = execution_price_per_share(c_hi, "buy", slippage)
        if None in (pay_fp, recv_np, pay_lo, recv_md, pay_hi):
            continue

        entry_net = (
            -float(pay_fp)
            + float(recv_np)
            - float(pay_lo)
            + 2.0 * float(recv_md)
            - float(pay_hi)
        ) * CONTRACT_MULTIPLIER

        exp_near = pd.Timestamp(r_np["expiration"]).normalize()
        exp_far = pd.Timestamp(r_fp["expiration"]).normalize()
        k_p = _scale_strike_raw(r_fp, spy_px)
        k_cl = _scale_strike_raw(r_lo, spy_px)
        k_cm = _scale_strike_raw(r_md, spy_px)
        k_ch = _scale_strike_raw(r_hi, spy_px)

        exit_ideal = _trading_day_offset(spy_idx, d, hold_days)
        last_near = _last_session_before_expiry(spy_idx, d, exp_near)
        if last_near is None or exit_ideal is None:
            continue
        exit_d = min(exit_ideal, last_near)

        pending = {
            "entry_d": d,
            "exit_d": exit_d,
            "entry_net_premium": entry_net,
            "spy_entry": spy_px,
            "exp_near": exp_near.date(),
            "exp_far": exp_far.date(),
            "k_p": k_p,
            "k_cl": k_cl,
            "k_cm": k_cm,
            "k_ch": k_ch,
            "pred_signal": pred_sig,
            "pred_far_put": parts["pred_far_put"],
            "pred_near_put": parts["pred_near_put"],
            "pred_c_lo": parts["pred_c_lo"],
            "pred_c_mid": parts["pred_c_mid"],
            "pred_c_hi": parts["pred_c_hi"],
            "c_far_put": c_fp,
            "c_near_put": c_np,
            "c_lo": c_lo,
            "c_mid": c_md,
            "c_hi": c_hi,
        }

    return trades


def main() -> None:
    ap = argparse.ArgumentParser(description="Flyagonal: put calendar + call butterfly (no SPY hedge)")
    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=0.0)
    ap.add_argument("--hold-days", type=int, default=5)
    ap.add_argument("--rebalance-every", type=int, default=5)
    ap.add_argument(
        "--dte-near-min",
        type=int,
        default=5,
        help="Front expiry DTE min (butterfly + short put)",
    )
    ap.add_argument("--dte-near-max", type=int, default=10)
    ap.add_argument(
        "--dte-far-min",
        type=int,
        default=12,
        help="Far put DTE min (should exceed near + gap)",
    )
    ap.add_argument("--dte-far-max", type=int, default=35)
    ap.add_argument("--calendar-gap-min", type=int, default=7)
    ap.add_argument("--calendar-gap-max", type=int, default=14)
    ap.add_argument("--put-moneyness-min", type=float, default=0.92)
    ap.add_argument("--put-moneyness-max", type=float, default=0.995)
    ap.add_argument(
        "--fly-strike-width",
        type=float,
        default=5.0,
        help="SPY points between adjacent butterfly strikes (lo-mid-hi)",
    )
    ap.add_argument("--min-pred-signal", type=float, default=-1e9)
    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,
        min_vix=float(args.min_vix),
        hold_days=int(args.hold_days),
        rebalance_every=int(args.rebalance_every),
        dte_near_min=int(args.dte_near_min),
        dte_near_max=int(args.dte_near_max),
        dte_far_min=int(args.dte_far_min),
        dte_far_max=int(args.dte_far_max),
        calendar_gap_min=int(args.calendar_gap_min),
        calendar_gap_max=int(args.calendar_gap_max),
        put_moneyness_min=float(args.put_moneyness_min),
        put_moneyness_max=float(args.put_moneyness_max),
        fly_strike_width=float(args.fly_strike_width),
        min_pred_signal=float(args.min_pred_signal),
        slippage=float(args.slippage),
    )

    pnls = [t.pnl_total for t in trades]
    print(
        json.dumps(
            {
                "n_trades": len(trades),
                "structure": "flyagonal",
                "total_pnl": float(np.sum(pnls)) if pnls else 0.0,
            },
            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()
