#!/usr/bin/env python3
"""
Always-on SPY skew regime book: steep → short RR, flat → long straddle, inverted → long RR.

Uses Theta 15:45 chains + same bid/ask rules as literature research.
One position at a time; on exit, next session enters the trade for that day's regime.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_skew_regime_always_on.py \\
      --start 2016-01-04 --end 2026-04-02 --capital 100000 \\
      --out-prefix RenTech/data/logs/skew_regime_always_on

    # Improvement sweep (single Theta precompute, all variants):
    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_skew_regime_always_on.py \\
      --sweep --start 2016-01-04 --end 2026-04-02 --capital 100000 \\
      --out-prefix RenTech/data/logs/skew_regime_variant_sweep
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Literal

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.options_backtest import find_contract_in_chain
from RenTech.core.options_data_loader import OptionChain
from RenTech.strategy_stack import research_literature_theta_strategies as L

Regime = Literal["steep", "flat", "inverted", "unknown"]

DTE = 40
PUT_D = -0.18
CALL_D = 0.12
STEEP_THR = 0.035
FLAT_LO = 0.0
STRADDLE_DTE = 18
HOLD = 5

S057_DTE = 30
S057_PUT_D = -0.20
S057_CALL_D = 0.12
S057_SKEW_THR = 0.056
S057_HOLD = 8


@dataclass(frozen=True)
class RegimeTrade:
    regime: Regime
    trade_kind: str
    trade_params: tuple
    label: str
    hold: int = HOLD


@dataclass(frozen=True)
class VariantSpec:
    name: str
    description: str
    steep: RegimeTrade | None
    flat: RegimeTrade | None
    inverted: RegimeTrade | None
    inverted_min_skew: float | None = None
    steep_s057_gates: bool = False


BASELINE_STEEP = RegimeTrade(
    "steep", "rr", (DTE, PUT_D, CALL_D), "Short RR (sell put, buy call)", HOLD
)
BASELINE_FLAT = RegimeTrade(
    "flat", "sl", (STRADDLE_DTE,), f"Long straddle ~{STRADDLE_DTE}D", HOLD
)
BASELINE_INVERTED = RegimeTrade(
    "inverted", "lrr", (DTE, PUT_D, CALL_D), "Long RR (buy put, sell call)", HOLD
)
S057_STEEP = RegimeTrade(
    "steep", "rr", (S057_DTE, S057_PUT_D, S057_CALL_D), "S057 short RR", S057_HOLD
)
INV_STRADDLE = RegimeTrade(
    "inverted", "sl", (STRADDLE_DTE,), f"Long straddle ~{STRADDLE_DTE}D", HOLD
)
INV_PUT_SPREAD = RegimeTrade(
    "inverted", "lvp", (30, -0.25, 7.0), "Long put debit spread", 10
)


def build_variants() -> list[VariantSpec]:
    return [
        VariantSpec(
            "baseline",
            "Original: steep short RR, flat straddle, inverted long RR",
            BASELINE_STEEP,
            BASELINE_FLAT,
            BASELINE_INVERTED,
        ),
        VariantSpec(
            "steep_only",
            "Steep short RR only; cash on flat/inverted",
            BASELINE_STEEP,
            None,
            None,
        ),
        VariantSpec(
            "s057_gated",
            "S057 short RR (30D, gates) only; cash otherwise",
            S057_STEEP,
            None,
            None,
            steep_s057_gates=True,
        ),
        VariantSpec(
            "inv_cash",
            "Baseline steep+flat; cash when inverted",
            BASELINE_STEEP,
            BASELINE_FLAT,
            None,
        ),
        VariantSpec(
            "inv_straddle",
            "Baseline steep+flat; inverted → long straddle",
            BASELINE_STEEP,
            BASELINE_FLAT,
            INV_STRADDLE,
        ),
        VariantSpec(
            "inv_putspread",
            "Baseline steep+flat; inverted → long put debit spread",
            BASELINE_STEEP,
            BASELINE_FLAT,
            INV_PUT_SPREAD,
        ),
        VariantSpec(
            "deep_inv_lrr",
            "Baseline steep+flat; long RR only when skew < -10%",
            BASELINE_STEEP,
            BASELINE_FLAT,
            BASELINE_INVERTED,
            inverted_min_skew=-0.10,
        ),
        VariantSpec(
            "s057_inv_cash",
            "S057 steep + cash on flat/inverted",
            S057_STEEP,
            None,
            None,
            steep_s057_gates=True,
        ),
    ]


REGIME_TRADES: dict[Regime, RegimeTrade] = {
    "steep": BASELINE_STEEP,
    "flat": BASELINE_FLAT,
    "inverted": BASELINE_INVERTED,
}


def classify_skew(skew: float | None) -> Regime:
    if skew is None or not math.isfinite(float(skew)):
        return "unknown"
    s = float(skew)
    if s > STEEP_THR:
        return "steep"
    if s < FLAT_LO:
        return "inverted"
    return "flat"


def long_risk_reversal_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    target_dte: int,
    put_delta: float,
    call_delta: float,
) -> float | None:
    p0 = L.find_target_leg_safe(entry_chain, target_dte, put_delta, "P")
    c0 = L.find_target_leg_safe(entry_chain, target_dte, call_delta, "C")
    if p0 is None or c0 is None or L._norm(c0.expiration) != L._norm(p0.expiration):
        return None
    exp = L._norm(p0.expiration)
    cost = (L._long_open_px(p0) - L._short_open_px(c0)) * L.MULT
    p_x = find_contract_in_chain(exit_chain, exp, float(p0.strike), "P")
    c_x = find_contract_in_chain(exit_chain, exp, float(c0.strike), "C")
    exit_d = L._norm(exit_chain.as_of)
    if p_x is not None and c_x is not None:
        val = (L._long_close_px(p_x) - L._short_close_px(c_x)) * L.MULT
        return float(val - cost)
    if exit_d >= exp:
        val = (L._settle_long(p0, exit_spy) - L._settle_short(c0, exit_spy)) * L.MULT
        return float(val - cost)
    return None


def long_vertical_put_pnl(
    entry_chain: OptionChain,
    exit_chain: OptionChain,
    exit_spy: float,
    target_dte: int,
    long_delta: float,
    wing_width: float,
) -> float | None:
    pl = L.find_target_leg_safe(entry_chain, target_dte, long_delta, "P")
    if pl is None:
        return None
    exp = L._norm(pl.expiration)
    short_strike = float(pl.strike) - float(wing_width)
    ps = find_contract_in_chain(entry_chain, exp, short_strike, "P")
    if ps is None:
        return None
    debit = (L._long_open_px(pl) - L._short_open_px(ps)) * L.MULT
    pl_x = find_contract_in_chain(exit_chain, exp, float(pl.strike), "P")
    ps_x = find_contract_in_chain(exit_chain, exp, short_strike, "P")
    exit_d = L._norm(exit_chain.as_of)
    if pl_x is not None and ps_x is not None:
        val = (L._long_close_px(pl_x) - L._short_close_px(ps_x)) * L.MULT
        return float(val - debit)
    if exit_d >= exp:
        val = (L._settle_long(pl, exit_spy) - L._settle_short(ps, exit_spy)) * L.MULT
        return float(val - debit)
    return None


def run_trade(
    kind: str,
    ch0: OptionChain,
    ch1: OptionChain,
    se: float,
    sx: float,
    tp: tuple,
) -> float | None:
    if kind == "rr":
        return L.short_risk_reversal_pnl(ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2]))
    if kind == "lrr":
        return long_risk_reversal_pnl(ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2]))
    if kind == "sl":
        return L.long_straddle_pnl(ch0, ch1, se, sx, int(tp[0]))
    if kind == "lvp":
        return long_vertical_put_pnl(ch0, ch1, sx, int(tp[0]), float(tp[1]), float(tp[2]))
    return None


def _s057_steep_ok(i: int, row: pd.Series, skew_map: dict[tuple[int, int, float, float], float | None]) -> bool:
    sk = skew_map.get((i, S057_DTE, S057_PUT_D, S057_CALL_D))
    if sk is None or float(sk) <= S057_SKEW_THR:
        return False
    s200 = row.get("sma_200")
    if s200 is None or pd.isna(s200) or float(row["close"]) <= float(s200):
        return False
    vx = row.get("vix_close")
    if vx is None or pd.isna(vx) or float(vx) >= 24.0:
        return False
    return True


def _pick_trade(reg: Regime, spec: VariantSpec) -> RegimeTrade | None:
    if reg == "steep":
        return spec.steep
    if reg == "flat":
        return spec.flat
    if reg == "inverted":
        return spec.inverted
    return None


def run_always_on(
    days: list[pd.Timestamp],
    panel: pd.DataFrame,
    get_chain,
    skew_map: dict[tuple[int, int, float, float], float | None],
    spec: VariantSpec,
) -> tuple[list[dict[str, Any]], pd.DataFrame]:
    trades: list[dict[str, Any]] = []
    regime_days: list[dict[str, Any]] = []
    pending: tuple[int, float, RegimeTrade, float, int] | None = None

    for i, d in enumerate(days):
        row = panel.reindex([L._norm(d)]).iloc[0]
        spy_c = float(row["close"])
        skew = skew_map.get((i, DTE, PUT_D, CALL_D))
        reg = classify_skew(skew)
        regime_days.append(
            {
                "date": L._norm(d).strftime("%Y-%m-%d"),
                "skew_iv_diff": skew,
                "regime": reg,
                "in_trade": pending is not None,
            }
        )

        ch = get_chain(d)
        if pending is not None:
            ent_i, spy_e, rt, skew_e, hold = pending
            if i >= ent_i + hold:
                ch_e = get_chain(days[ent_i])
                pnl = run_trade(rt.trade_kind, ch_e, ch, spy_e, spy_c, rt.trade_params)
                pending = None
                if pnl is not None and math.isfinite(pnl):
                    trades.append(
                        {
                            "variant": spec.name,
                            "regime": rt.regime,
                            "trade_kind": rt.trade_kind,
                            "label": rt.label,
                            "entry_date": L._norm(days[ent_i]).strftime("%Y-%m-%d"),
                            "exit_date": L._norm(d).strftime("%Y-%m-%d"),
                            "skew_entry": skew_e,
                            "pnl_usd": float(pnl),
                        }
                    )
        if pending is not None:
            continue
        if reg == "unknown" or not ch.contracts:
            continue
        rt = _pick_trade(reg, spec)
        if rt is None:
            continue
        if spec.steep_s057_gates and reg == "steep" and not _s057_steep_ok(i, row, skew_map):
            continue
        if (
            reg == "inverted"
            and spec.inverted_min_skew is not None
            and (skew is None or float(skew) >= float(spec.inverted_min_skew))
        ):
            continue
        pending = (i, spy_c, rt, float(skew) if skew is not None else float("nan"), int(rt.hold))

    daily = pd.DataFrame(regime_days)
    daily["date"] = pd.to_datetime(daily["date"])
    daily = daily.set_index("date").sort_index()
    return trades, daily


def equity_metrics(eq: pd.Series, cap: float) -> dict[str, float]:
    eq = eq.astype(float)
    ret = eq.pct_change().fillna(0.0)
    sh = (
        float(ret.mean()) / float(ret.std(ddof=1)) * math.sqrt(252)
        if float(ret.std(ddof=1)) > 1e-12
        else 0.0
    )
    dd = float((eq / eq.cummax() - 1.0).min())
    tot = (float(eq.iloc[-1]) / cap - 1.0) * 100.0
    return {"return_pct": tot, "sharpe": sh, "max_dd_pct": dd * 100.0, "end_equity": float(eq.iloc[-1])}


def yearly_returns(eq: pd.Series, cap: float) -> dict[int, float]:
    out: dict[int, float] = {}
    for y in sorted(eq.index.year.unique()):
        prior = eq.loc[eq.index.year < y]
        sub = eq.loc[eq.index.year == y]
        if sub.empty:
            continue
        start = float(prior.iloc[-1]) if len(prior) else cap
        end = float(sub.iloc[-1])
        out[int(y)] = (end / start - 1.0) * 100.0
    return out


def summarize_variant(
    spec: VariantSpec,
    trades: list[dict[str, Any]],
    daily_reg: pd.DataFrame,
    cap: float,
) -> dict[str, Any]:
    pnl_by_day = pd.Series(0.0, index=daily_reg.index)
    for t in trades:
        ex = pd.Timestamp(t["exit_date"])
        if ex in pnl_by_day.index:
            pnl_by_day.loc[ex] += float(t["pnl_usd"])
    eq = cap + pnl_by_day.cumsum()

    by_reg: dict[str, list[float]] = {k: [] for k in ("steep", "flat", "inverted")}
    for t in trades:
        by_reg[t["regime"]].append(float(t["pnl_usd"]))

    in_trade_pct = float(daily_reg["in_trade"].mean() * 100.0)
    return {
        "name": spec.name,
        "description": spec.description,
        "trades_by_regime": {
            k: {
                "n_trades": len(v),
                "sum_pnl_usd": round(sum(v), 2),
                "avg_pnl_usd": round(float(np.mean(v)), 2) if v else 0.0,
                "win_rate_pct": round(100.0 * sum(1 for x in v if x > 0) / len(v), 1) if v else 0.0,
            }
            for k, v in by_reg.items()
        },
        "in_trade_pct": round(in_trade_pct, 1),
        "combined": equity_metrics(eq, cap),
        "yearly_return_pct": yearly_returns(eq, cap),
        "n_trades_total": len(trades),
        "equity": eq,
        "daily_pnl": pnl_by_day,
        "daily_reg": daily_reg,
    }


def _resolve_prefix(prefix: Path) -> Path:
    p = prefix.expanduser()
    if p.suffix:
        p = p.with_suffix("")
    p.parent.mkdir(parents=True, exist_ok=True)
    return p


def write_variant_artifacts(prefix: Path, summary: dict[str, Any], trades: list[dict[str, Any]]) -> None:
    trades_path = Path(str(prefix) + "_trades.csv")
    daily_path = Path(str(prefix) + "_daily.csv")
    meta_path = Path(str(prefix) + "_meta.json")
    pd.DataFrame(trades).to_csv(trades_path, index=False)
    out_daily = summary["daily_reg"].copy()
    out_daily["daily_pnl_usd"] = summary["daily_pnl"]
    out_daily["equity_usd"] = summary["equity"]
    out_daily.to_csv(daily_path)
    meta = {k: v for k, v in summary.items() if k not in ("equity", "daily_pnl", "daily_reg")}
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")


def print_variant_summary(summary: dict[str, Any]) -> None:
    c = summary["combined"]
    print(
        f"  {summary['name']}: return={c['return_pct']:.1f}%  Sharpe={c['sharpe']:.2f}  "
        f"maxDD={c['max_dd_pct']:.2f}%  trades={summary['n_trades_total']}  "
        f"in_trade={summary['in_trade_pct']:.1f}%",
        flush=True,
    )


def run_sweep(
    *,
    days: list[pd.Timestamp],
    panel: pd.DataFrame,
    get_chain,
    skew_map: dict,
    cap: float,
    out_prefix: Path,
    variants: list[VariantSpec],
) -> list[dict[str, Any]]:
    results: list[dict[str, Any]] = []
    print("\n--- Variant sweep ---", flush=True)
    for spec in variants:
        trades, daily_reg = run_always_on(days, panel, get_chain, skew_map, spec)
        summary = summarize_variant(spec, trades, daily_reg, cap)
        print_variant_summary(summary)
        results.append(summary)
        vprefix = _resolve_prefix(out_prefix.parent / f"{out_prefix.name}_{spec.name}")
        write_variant_artifacts(vprefix, summary, trades)

    ranked = sorted(results, key=lambda r: r["combined"]["sharpe"], reverse=True)
    sweep_meta = {
        "variants": [{k: v for k, v in r.items() if k not in ("equity", "daily_pnl", "daily_reg")} for r in results],
        "ranked_by_sharpe": [r["name"] for r in ranked],
    }
    sweep_path = _resolve_prefix(out_prefix)
    Path(str(sweep_path) + "_summary.json").write_text(
        json.dumps(sweep_meta, indent=2) + "\n", encoding="utf-8"
    )

    md_lines = [
        "# Skew regime variant sweep\n",
        "| Variant | Return | Sharpe | Max DD | Trades | In-trade % |\n",
        "|---------|--------|--------|--------|--------|------------|\n",
    ]
    for r in ranked:
        c = r["combined"]
        md_lines.append(
            f"| {r['name']} | {c['return_pct']:.1f}% | {c['sharpe']:.2f} | "
            f"{c['max_dd_pct']:.2f}% | {r['n_trades_total']} | {r['in_trade_pct']:.1f}% |\n"
        )
    Path(str(sweep_path) + "_summary.md").write_text("".join(md_lines), encoding="utf-8")
    print(f"\nWrote sweep summary to {sweep_path}_summary.json", flush=True)
    return results


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-04-02")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--hold", type=int, default=HOLD, help="Default hold (baseline single-variant mode only)")
    ap.add_argument("--variant", default="baseline", help="Variant name when not using --sweep")
    ap.add_argument("--sweep", action="store_true", help="Run all improvement variants after one precompute")
    ap.add_argument("--out-prefix", type=Path, default=_REPO / "RenTech/data/logs/skew_regime_always_on")
    args = ap.parse_args()

    cap = float(args.capital)
    print(f"Skew regime always-on · {args.start} → {args.end}", flush=True)
    days, panel, get_chain, iv_atm, skew_map, n_contracts, _ = L.prepare_theta_research_context(
        theta_dir=L._DEFAULT_THETA,
        capital=cap,
        start=str(args.start),
        end=str(args.end),
        max_days=0,
    )
    print(f"  sessions={len(days)}", flush=True)

    if args.sweep:
        run_sweep(
            days=days,
            panel=panel,
            get_chain=get_chain,
            skew_map=skew_map,
            cap=cap,
            out_prefix=_resolve_prefix(args.out_prefix),
            variants=build_variants(),
        )
        return

    variants = {v.name: v for v in build_variants()}
    spec = variants.get(str(args.variant), variants["baseline"])
    if str(args.variant) not in variants:
        print(f"Unknown variant {args.variant!r}; using baseline.", flush=True)

    trades, daily_reg = run_always_on(days, panel, get_chain, skew_map, spec)
    summary = summarize_variant(spec, trades, daily_reg, cap)
    prefix = _resolve_prefix(args.out_prefix)
    write_variant_artifacts(prefix, summary, trades)

    regime_share = daily_reg["regime"].value_counts(normalize=True).mul(100).round(1).to_dict()
    print("\n--- Regime day mix (% of sessions) ---", flush=True)
    for k, v in sorted(regime_share.items()):
        print(f"  {k}: {v}%", flush=True)
    print(f"  in_trade: {summary['in_trade_pct']:.1f}%", flush=True)

    print("\n--- Trades by regime ---", flush=True)
    for k, st in summary["trades_by_regime"].items():
        print(
            f"  {k}: n={st['n_trades']}  sum_pnl=${st['sum_pnl_usd']:,.0f}  "
            f"avg=${st['avg_pnl_usd']:,.0f}  win={st['win_rate_pct']:.1f}%",
            flush=True,
        )

    print("\n--- Combined ---", flush=True)
    print_variant_summary(summary)
    print(f"\nWrote {prefix}_{{trades,daily,meta}}.*", flush=True)


if __name__ == "__main__":
    main()
