#!/usr/bin/env python3
"""
SPX/SPY **Regime State Space** strategy (v2 spec) — clean standalone build.

Six features → regime intersections (A/B/C/D) → options action mapping.
Uses Theta 15:45 SPY chains for IV/skew and execution; Yahoo for macro panel.

Modes
-----
* ``--mode states`` — daily feature + regime report only (no option trades).
* ``--mode backtest`` — regime-router always-on book (default).

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_spx_regime_state_space.py \\
        --start 2016-01-04 --end 2022-12-30 --capital 100000 \\
        --out-prefix RenTech/data/logs/spx_regime_state_space

    # Fast diagnostic (no Theta execution):
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/run_spx_regime_state_space.py \\
        --mode states --start 2016-01-04 --end 2022-12-30 \\
        --out-prefix RenTech/data/logs/spx_regime_state_space_states
"""

from __future__ import annotations

import argparse
import json
import math
import sys
import time
from pathlib import Path
from typing import Any

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, theta_chunks_date_bounds
from RenTech.strategy_stack import research_literature_theta_strategies as L
from RenTech.strategy_stack.spx_regime_state_space.backtest import build_states_only, run_regime_router
from RenTech.strategy_stack.spx_regime_state_space.features import (
    IV_TARGET_DTE,
    SKEW_CALL_DELTA,
    SKEW_PUT_DELTA,
    build_regime_features,
)
from RenTech.strategy_stack.vrp_backtester import (
    load_spy_vix_from_yfinance,
    normalize_spy_df,
    trading_days_intersecting_spy,
)

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT = LOGS / "spx_regime_state_space"
DEFAULT_THETA = _REPO / "RenTech" / "data" / "theta_chunks"


def _precompute_iv_skew(
    days: list[pd.Timestamp],
    get_chain,
    spy_close: pd.Series,
) -> tuple[dict[pd.Timestamp, float | None], dict[pd.Timestamp, float | None]]:
    iv_by_day: dict[pd.Timestamp, float | None] = {}
    skew_by_day: dict[pd.Timestamp, float | None] = {}
    for i, d in enumerate(days):
        dt = L._norm(d)
        ch = get_chain(d)
        row_close = float(spy_close.reindex([dt]).iloc[0])
        iv_by_day[dt] = L.atm_iv_straddle(ch, row_close, IV_TARGET_DTE)
        pleg = L.find_target_leg_safe(ch, IV_TARGET_DTE, SKEW_PUT_DELTA, "P")
        cleg = L.find_target_leg_safe(ch, IV_TARGET_DTE, SKEW_CALL_DELTA, "C")
        if (
            pleg is not None
            and cleg is not None
            and math.isfinite(float(pleg.iv))
            and math.isfinite(float(cleg.iv))
        ):
            skew_by_day[dt] = float(pleg.iv) - float(cleg.iv)
        else:
            skew_by_day[dt] = None
        if (i + 1) % 250 == 0:
            print(f"  IV/skew cache {i + 1}/{len(days)}", flush=True)
    return iv_by_day, skew_by_day


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
    n = len(ret)
    years = n / 252.0 if n else 0.0
    cagr = ((float(eq.iloc[-1]) / cap) ** (1.0 / years) - 1.0) * 100.0 if years > 0 else 0.0
    return {
        "return_pct": tot,
        "cagr_pct": cagr,
        "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 _regime_mix(daily: pd.DataFrame) -> dict[str, float]:
    if "regime" not in daily.columns or daily.empty:
        return {}
    return daily["regime"].value_counts(normalize=True).mul(100).round(1).to_dict()


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--mode", choices=("backtest", "states"), default="backtest")
    ap.add_argument("--start", default="2016-01-04")
    ap.add_argument("--end", default="")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--hold", type=int, default=5, help="Sessions to hold each structure")
    ap.add_argument("--theta-dir", type=Path, default=DEFAULT_THETA)
    ap.add_argument("--max-days", type=int, default=0, help="Smoke: cap session count")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    cap = float(args.capital)
    theta_dir = args.theta_dir.expanduser().resolve()
    t0 = time.perf_counter()

    d0, d1 = theta_chunks_date_bounds(theta_dir)
    yf_start = (d0 - pd.Timedelta(days=500)).strftime("%Y-%m-%d")
    yf_end = (d1 + pd.Timedelta(days=14)).strftime("%Y-%m-%d")
    spy_wide = normalize_spy_df(load_spy_vix_from_yfinance(yf_start, yf_end))
    ld = ThetaChunksLoader(theta_dir, spy_wide)
    days = trading_days_intersecting_spy(ld, spy_wide.index, d0, d1)
    if str(args.start).strip():
        t_start = pd.Timestamp(str(args.start).strip())
        days = [d for d in days if L._norm(d) >= t_start]
    if str(args.end).strip():
        t_end = pd.Timestamp(str(args.end).strip())
        days = [d for d in days if L._norm(d) <= t_end]
    if int(args.max_days) > 0:
        days = days[: int(args.max_days)]

    if not days:
        raise SystemExit("No trading sessions in window (check --start/--end and Theta data).")

    print(f"SPX regime state space · mode={args.mode} · {args.start} → {args.end or 'data end'}", flush=True)
    print(f"  sessions={len(days)}", flush=True)

    chain_cache: dict[pd.Timestamp, Any] = {}
    _CHAIN_KW = dict(strike_pct_lo=0.76, strike_pct_hi=1.12, min_dte=5, max_dte=60)

    def get_chain(d: pd.Timestamp):
        t = L._norm(d)
        if t not in chain_cache:
            chain_cache[t] = ld.get_chain_for_date(t, **_CHAIN_KW)
        return chain_cache[t]

    print("Precomputing IV30 + 25Δ skew from Theta chains …", flush=True)
    iv_by_day, skew_by_day = _precompute_iv_skew(days, get_chain, spy_wide["close"])

    features = build_regime_features(
        spy_wide,
        days,
        iv30_by_day=iv_by_day,
        skew25_by_day=skew_by_day,
    )
    features.index = [L._norm(d) for d in features.index]

    prefix = args.out_prefix.expanduser()
    if prefix.suffix:
        prefix = prefix.with_suffix("")
    prefix.parent.mkdir(parents=True, exist_ok=True)

    meta: dict[str, Any] = {
        "spec": "SPX_Regime_State_Space_v2",
        "mode": args.mode,
        "window": {
            "start": str(args.start),
            "end": str(args.end) or None,
            "sessions": len(days),
            "first": L._norm(days[0]).strftime("%Y-%m-%d") if days else None,
            "last": L._norm(days[-1]).strftime("%Y-%m-%d") if days else None,
        },
        "capital_usd": cap,
        "hold_sessions": int(args.hold),
        "regime_day_pct": {},
        "action_day_pct": {},
        "combined": {},
        "trades_by_regime": {},
        "n_trades_total": 0,
    }

    if args.mode == "states":
        daily = build_states_only(features)
        meta["regime_day_pct"] = _regime_mix(daily)
        meta["action_day_pct"] = (
            daily["action"].value_counts(normalize=True).mul(100).round(1).to_dict()
            if not daily.empty
            else {}
        )
        daily_path = Path(str(prefix) + "_states_daily.csv")
        daily.reset_index().to_csv(daily_path, index=False)
        meta_path = Path(str(prefix) + "_states_meta.json")
        meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")
        print(f"\nRegime mix (% sessions):", flush=True)
        for k, v in sorted(meta["regime_day_pct"].items()):
            print(f"  {k}: {v}%", flush=True)
        print(f"\nWrote {daily_path}\n      {meta_path} ({time.perf_counter() - t0:.1f}s)", flush=True)
        return

    trades, daily = run_regime_router(
        days,
        features,
        get_chain,
        hold=int(args.hold),
    )

    pnl_by_day = pd.Series(0.0, index=daily.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_regime: dict[str, list[float]] = {}
    for t in trades:
        by_regime.setdefault(t["regime"], []).append(float(t["pnl_usd"]))

    meta["regime_day_pct"] = _regime_mix(daily)
    meta["action_day_pct"] = daily["action"].value_counts(normalize=True).mul(100).round(1).to_dict()
    meta["in_trade_pct"] = round(float(daily["in_trade"].mean() * 100.0), 1)
    meta["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_regime.items()
    }
    meta["combined"] = equity_metrics(eq, cap)
    meta["yearly_return_pct"] = yearly_returns(eq, cap)
    meta["n_trades_total"] = len(trades)

    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 = daily.copy()
    out_daily["daily_pnl_usd"] = pnl_by_day
    out_daily["equity_usd"] = eq
    out_daily.reset_index().to_csv(daily_path, index=False)
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")

    print("\n--- Regime day mix (% of sessions) ---", flush=True)
    for k, v in sorted(meta["regime_day_pct"].items()):
        print(f"  {k}: {v}%", flush=True)
    print(f"  in_trade: {meta['in_trade_pct']}%", flush=True)

    print("\n--- Trades by regime ---", flush=True)
    for k, st in sorted(meta["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,
        )

    c = meta["combined"]
    print(
        f"\nCombined: return={c['return_pct']:.1f}%  CAGR={c['cagr_pct']:.1f}%  "
        f"Sharpe={c['sharpe']:.2f}  maxDD={c['max_dd_pct']:.2f}%  "
        f"end=${c['end_equity']:,.0f}  trades={len(trades)}",
        flush=True,
    )
    print(f"\nWrote {trades_path}\n      {daily_path}\n      {meta_path} ({time.perf_counter() - t0:.1f}s)", flush=True)


if __name__ == "__main__":
    main()
