#!/usr/bin/env python3
"""
Markov-chain backtest on **VIX calendar spreads** (M1−M2, M2−M3, M3−M4, M4−M5).

Each spread is modeled as its own Markov state machine on the **rolling percentile**
of the spread level in VIX points (e.g. VX1=18, VX2=20 → spread = −2).

**Mispricing logic:**
  - ``market_price`` = current spread percentile (where the spread sits in its range)
  - Monte Carlo → ``calibrated_prob`` of ending in the upper half of spread states
  - ``edge = calibrated_prob − market_price``
  - edge > threshold  → **LONG spread** (buy near, sell far — spread widens)
  - edge < −threshold → **SHORT spread** (sell near, buy far — spread narrows)
  - |edge| ≤ threshold → flat (cash)

**Execution modes:**
  - ``--fixed-expiry`` (default): locks specific ``near_expiry`` / ``far_expiry`` at entry;
    rolls explicitly on near-leg expiry via ``vix_fixed_calendar_engine``.
  - ``--no-fixed-expiry``: constant-tenor rolled series (research only — trade log
    misstates long holds across contract rolls).

**PnL (fixed-expiry):** daily return on locked calendar ≈ ``ret(near_locked) − ret(far_locked)``;
scaled by signed Kelly exposure × per-spread budget.

Data: ``RenTech/data/vix_futures_cboe.parquet`` + ``RenTech/data/vix_futures_contracts_long.parquet``
(run ``download_cboe_vix_futures.py``).

Example::

    cd /Users/robzingale/trading_bot
    .venv/bin/python RenTech/data_pipeline/download_cboe_vix_futures.py
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_markov_vix_spread_backtest.py \\
        --start 2016-01-04
"""

from __future__ import annotations

import argparse
import json
import subprocess
import sys
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.strategy_stack.markov_chain_trading import (
    MarkovChainTradingModel,
    walk_forward_exposure_series,
)

from RenTech.strategy_stack.vix_fixed_calendar_engine import (
    VixContractStore,
    simulate_fixed_calendar_portfolio,
)

DATA_DIR = _REPO / "RenTech" / "data"
LOGS = DATA_DIR / "logs"
VIX_FUT_PATH = DATA_DIR / "vix_futures_cboe.parquet"
CONTRACTS_PATH = DATA_DIR / "vix_futures_contracts_long.parquet"
DEFAULT_OUT_PREFIX = LOGS / "markov_vix_spread"

# (label, spread_col, near_vx_col, far_vx_col)
SPREAD_LEGS: list[tuple[str, str, str, str]] = [
    ("M1_M2", "spread_m1_m2", "vx1", "vx2"),
    ("M2_M3", "spread_m2_m3", "vx2", "vx3"),
    ("M3_M4", "spread_m3_m4", "vx3", "vx4"),
    ("M4_M5", "spread_m4_m5", "vx4", "vx5"),
]


def _ensure_panel(force_download: bool = False) -> pd.DataFrame:
    if force_download or not VIX_FUT_PATH.is_file():
        script = _REPO / "RenTech" / "data_pipeline" / "download_cboe_vix_futures.py"
        subprocess.run([sys.executable, str(script)], cwd=str(_REPO), check=True)
    panel = pd.read_parquet(VIX_FUT_PATH)
    panel.index = pd.to_datetime(panel.index).tz_localize(None).sort_values()
    # Rebuild spreads if parquet predates vx4/vx5 columns
    for n in range(1, 6):
        col = f"vx{n}"
        if col not in panel.columns:
            settle = panel.get(f"{col}_settle")
            close = panel.get(f"{col}_close")
            if settle is not None or close is not None:
                panel[col] = settle.fillna(close) if settle is not None else close
    if "spread_m1_m2" not in panel.columns and "vx1" in panel.columns:
        panel["spread_m1_m2"] = panel["vx1"] - panel["vx2"]
        panel["spread_m2_m3"] = panel["vx2"] - panel["vx3"]
        panel["spread_m3_m4"] = panel["vx3"] - panel["vx4"]
        panel["spread_m4_m5"] = panel["vx4"] - panel["vx5"]
    return panel


def _metrics(daily_ret: pd.Series, capital: float) -> dict[str, float]:
    r = daily_ret.dropna().astype(np.float64)
    n = len(r)
    if n == 0:
        return {
            "total_return_pct": float("nan"),
            "cagr_pct": float("nan"),
            "sharpe": float("nan"),
            "max_drawdown_pct": float("nan"),
            "ending_equity_usd": capital,
        }
    eq = capital * (1.0 + r).cumprod()
    years = n / 252.0
    end_eq = float(eq.iloc[-1])
    total_ret = end_eq / capital - 1.0
    cagr = (end_eq / capital) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = eq / eq.cummax() - 1.0
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    return {
        "total_return_pct": float(total_ret * 100.0),
        "cagr_pct": float(cagr * 100.0),
        "sharpe": sharpe,
        "max_drawdown_pct": float(dd.min() * 100.0),
        "ending_equity_usd": end_eq,
    }


def build_spread_signals(
    panel: pd.DataFrame,
    *,
    model: MarkovChainTradingModel,
    spreads: list[tuple[str, str, str, str]],
    percentile_window: int,
    lookback: int,
    matrix_refresh: int,
    allow_short: bool,
    max_short: float,
    gross_cap: float,
    vix: pd.Series | None = None,
    vix_lo: float = 15.0,
    vix_hi: float = 25.0,
) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame, pd.DataFrame, list[str]]:
    """Markov signals on constant-tenor spread levels (for mispricing detection only)."""
    master = panel.index.sort_values()
    n_spreads = len(spreads)
    budget = 1.0 / n_spreads

    weight_df = pd.DataFrame(0.0, index=master, columns=[s[0] for s in spreads])
    spread_ret_df = pd.DataFrame(0.0, index=master, columns=[s[0] for s in spreads])
    edge_df = pd.DataFrame(np.nan, index=master, columns=[s[0] for s in spreads])
    level_df = pd.DataFrame(np.nan, index=master, columns=[s[0] for s in spreads])
    decision_df = pd.DataFrame("PASS", index=master, columns=[s[0] for s in spreads])

    vix_aligned = vix.reindex(master).ffill() if vix is not None else None

    for label, spread_col, near_col, far_col in spreads:
        if spread_col not in panel.columns:
            print(f"  Skip {label}: missing {spread_col}", flush=True)
            continue
        near = panel[near_col].astype(np.float64).reindex(master).ffill()
        far = panel[far_col].astype(np.float64).reindex(master).ffill()
        spread_level = panel[spread_col].astype(np.float64).reindex(master).ffill()
        level_df[label] = spread_level
        spread_ret_df[label] = near.pct_change().fillna(0.0) - far.pct_change().fillna(0.0)

        sig = walk_forward_exposure_series(
            spread_level,
            percentile_window=percentile_window,
            lookback=lookback,
            model=model,
            matrix_refresh=matrix_refresh,
            vix=vix_aligned,
            vix_lo=vix_lo,
            vix_hi=vix_hi,
            allow_short=allow_short,
            max_short_exposure=max_short,
        )
        if sig.empty:
            continue
        w = (sig["exposure"] * budget).reindex(master).fillna(0.0)
        w = w.clip(lower=-max_short * budget, upper=budget)
        weight_df[label] = w
        edge_df[label] = sig["edge"].reindex(master)
        decision_df[label] = sig["decision"].reindex(master).fillna("PASS")

    active = [c for c in weight_df.columns if c in spread_ret_df.columns]
    gross = weight_df[active].abs().sum(axis=1)
    over = gross > gross_cap
    if over.any():
        scale = (gross_cap / gross).clip(upper=1.0)
        weight_df.loc[over, active] = weight_df.loc[over, active].mul(scale[over], axis=0)

    return weight_df, edge_df, decision_df, level_df, active


def build_spread_portfolio(
    panel: pd.DataFrame,
    *,
    model: MarkovChainTradingModel,
    spreads: list[tuple[str, str, str, str]],
    percentile_window: int,
    lookback: int,
    cash_annual_yield: float,
    matrix_refresh: int,
    allow_short: bool,
    max_short: float,
    gross_cap: float,
    vix: pd.Series | None = None,
    vix_lo: float = 15.0,
    vix_hi: float = 25.0,
) -> pd.DataFrame:
    master = panel.index.sort_values()
    n_spreads = len(spreads)
    budget = 1.0 / n_spreads

    weight_df = pd.DataFrame(0.0, index=master, columns=[s[0] for s in spreads])
    spread_ret_df = pd.DataFrame(0.0, index=master, columns=[s[0] for s in spreads])
    edge_df = pd.DataFrame(np.nan, index=master, columns=[s[0] for s in spreads])
    level_df = pd.DataFrame(np.nan, index=master, columns=[s[0] for s in spreads])
    decision_df = pd.DataFrame("PASS", index=master, columns=[s[0] for s in spreads])

    vix_aligned = vix.reindex(master).ffill() if vix is not None else None

    for label, spread_col, near_col, far_col in spreads:
        if spread_col not in panel.columns:
            print(f"  Skip {label}: missing {spread_col}", flush=True)
            continue
        near = panel[near_col].astype(np.float64).reindex(master).ffill()
        far = panel[far_col].astype(np.float64).reindex(master).ffill()
        spread_level = panel[spread_col].astype(np.float64).reindex(master).ffill()
        level_df[label] = spread_level

        # Calendar spread return: long near + short far
        spread_ret = near.pct_change().fillna(0.0) - far.pct_change().fillna(0.0)
        spread_ret_df[label] = spread_ret

        sig = walk_forward_exposure_series(
            spread_level,
            percentile_window=percentile_window,
            lookback=lookback,
            model=model,
            matrix_refresh=matrix_refresh,
            vix=vix_aligned,
            vix_lo=vix_lo,
            vix_hi=vix_hi,
            allow_short=allow_short,
            max_short_exposure=max_short,
        )
        if sig.empty:
            continue

        w = (sig["exposure"] * budget).reindex(master).fillna(0.0)
        w = w.clip(lower=-max_short * budget, upper=budget)
        weight_df[label] = w
        edge_df[label] = sig["edge"].reindex(master)
        decision_df[label] = sig["decision"].reindex(master).fillna("PASS")

    active = [c for c in weight_df.columns if spread_ret_df[c].abs().sum() > 0 or weight_df[c].abs().sum() > 0]
    if not active:
        raise RuntimeError("No spread series available for backtest")

    gross = weight_df[active].abs().sum(axis=1)
    over = gross > gross_cap
    if over.any():
        scale = (gross_cap / gross).clip(upper=1.0)
        weight_df.loc[over, active] = weight_df.loc[over, active].mul(scale[over], axis=0)
        gross = weight_df[active].abs().sum(axis=1)

    cash_weight = (1.0 - gross).clip(lower=0.0)
    daily_rf = cash_annual_yield / 252.0

    port_ret = (spread_ret_df[active] * weight_df[active]).sum(axis=1) + cash_weight * daily_rf

    # Naive benchmark: always short each spread (harvest contango — sell near, buy far)
    naive_w = pd.Series({c: -budget for c in active})
    naive_ret = spread_ret_df[active].mul(naive_w, axis=1).sum(axis=1)
    naive_gross = abs(-budget) * len(active)
    naive_cash = max(0.0, 1.0 - naive_gross)
    naive_ret = naive_ret + naive_cash * daily_rf

    out = pd.DataFrame(
        {
            "portfolio_bar_ret": port_ret,
            "naive_short_spread_ret": naive_ret,
            "cash_weight": cash_weight,
            "gross_exposure": gross,
        },
        index=master,
    )
    for label in active:
        out[f"weight_{label}"] = weight_df[label]
        out[f"edge_{label}"] = edge_df[label]
        out[f"spread_{label}"] = level_df[label]
        out[f"decision_{label}"] = decision_df[label]
    return out


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="")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--cash-yield", type=float, default=0.04)
    ap.add_argument("--n-states", type=int, default=10)
    ap.add_argument("--horizon", type=int, default=21)
    ap.add_argument("--n-sims", type=int, default=3000)
    ap.add_argument("--percentile-window", type=int, default=252)
    ap.add_argument("--lookback", type=int, default=252)
    ap.add_argument("--edge-threshold", type=float, default=0.03)
    ap.add_argument(
        "--calibration", choices=("none", "equity_shrink", "polymarket"), default="equity_shrink"
    )
    ap.add_argument("--kelly-mult", type=float, default=0.25)
    ap.add_argument("--matrix-refresh", type=int, default=5)
    ap.add_argument("--max-short", type=float, default=0.50)
    ap.add_argument("--gross-cap", type=float, default=1.50)
    ap.add_argument("--vix-regime", action="store_true")
    ap.add_argument("--vix-lo", type=float, default=15.0)
    ap.add_argument("--vix-hi", type=float, default=25.0)
    ap.add_argument("--force-download", action="store_true")
    ap.add_argument(
        "--fixed-expiry",
        action=argparse.BooleanOptionalAction,
        default=True,
        help="Execute on locked near/far contract expiries with explicit rolls (default)",
    )
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    args = ap.parse_args()

    model = MarkovChainTradingModel(
        n_states=args.n_states,
        horizon=args.horizon,
        n_sims=args.n_sims,
        kelly_mult=args.kelly_mult,
        edge_threshold=args.edge_threshold,
        calibration_mode=args.calibration,
    )

    panel = _ensure_panel(force_download=args.force_download)
    vix = panel["vix_spot"].astype(np.float64) if args.vix_regime else None

    # Only include spreads with sufficient data
    available: list[tuple[str, str, str, str]] = []
    for legs in SPREAD_LEGS:
        label, spread_col, near, far = legs
        if spread_col in panel.columns and panel[spread_col].notna().sum() > 500:
            available.append(legs)
        else:
            print(f"  {label}: insufficient data, skipping", flush=True)

    print(f"Spreads: {[s[0] for s in available]}", flush=True)
    print(f"Execution mode: {'fixed-expiry calendar' if args.fixed_expiry else 'constant-tenor rolled'}", flush=True)
    print("Building Markov spread signals …", flush=True)

    weight_df, edge_df, decision_df, level_df, active = build_spread_signals(
        panel,
        model=model,
        spreads=available,
        percentile_window=args.percentile_window,
        lookback=args.lookback,
        matrix_refresh=args.matrix_refresh,
        allow_short=True,
        max_short=args.max_short,
        gross_cap=args.gross_cap,
        vix=vix,
        vix_lo=args.vix_lo,
        vix_hi=args.vix_hi,
    )

    trade_log_df: pd.DataFrame | None = None

    if args.fixed_expiry:
        if not CONTRACTS_PATH.is_file():
            print(f"Building {CONTRACTS_PATH} …", flush=True)
            subprocess.run(
                [sys.executable, str(_REPO / "RenTech/data_pipeline/download_cboe_vix_futures.py")],
                cwd=str(_REPO),
                check=True,
            )
        store = VixContractStore(pd.read_parquet(CONTRACTS_PATH))
        master = weight_df.index.sort_values()
        port, trade_log_df = simulate_fixed_calendar_portfolio(
            store=store,
            dates=master,
            target_weights=weight_df[active],
            edge_df=edge_df[active],
            decision_df=decision_df[active],
            spreads=active,
            capital=float(args.capital),
            cash_annual_yield=float(args.cash_yield),
            gross_cap=args.gross_cap,
        )
        for label in active:
            port[f"spread_{label}"] = level_df[label]
        # Constant-tenor benchmark for comparison
        budget = 1.0 / len(active)
        spread_ret = pd.DataFrame(index=master)
        for label, spread_col, near_col, far_col in available:
            if label not in active:
                continue
            near = panel[near_col].astype(float).reindex(master).ffill()
            far = panel[far_col].astype(float).reindex(master).ffill()
            spread_ret[label] = near.pct_change().fillna(0.0) - far.pct_change().fillna(0.0)
        naive_w = pd.Series({c: -budget for c in active})
        daily_rf = float(args.cash_yield) / 252.0
        port["naive_short_spread_ret"] = spread_ret[active].mul(naive_w, axis=1).sum(axis=1) + daily_rf
        port["execution_mode"] = "fixed_expiry_calendar"
    else:
        port = build_spread_portfolio(
            panel,
            model=model,
            spreads=available,
            percentile_window=args.percentile_window,
            lookback=args.lookback,
            cash_annual_yield=float(args.cash_yield),
            matrix_refresh=args.matrix_refresh,
            allow_short=True,
            max_short=args.max_short,
            gross_cap=args.gross_cap,
            vix=vix,
            vix_lo=args.vix_lo,
            vix_hi=args.vix_hi,
        )
        port["execution_mode"] = "constant_tenor_rolled"

    mask = port.index >= pd.Timestamp(args.start)
    if args.end.strip():
        mask &= port.index <= pd.Timestamp(args.end)
    win = port.loc[mask]

    cap = float(args.capital)
    r_mk = win["portfolio_bar_ret"].fillna(0.0).astype(np.float64)
    r_naive = win["naive_short_spread_ret"].fillna(0.0).astype(np.float64)

    m_mk = _metrics(r_mk, cap)
    m_naive = _metrics(r_naive, cap)

    short_cols = [c for c in win.columns if c.startswith("weight_")]
    short_days = int((win[short_cols].values < -0.005).any(axis=1).sum()) if short_cols else 0
    long_days = int((win[short_cols].values > 0.005).any(axis=1).sum()) if short_cols else 0

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    meta_path = Path(f"{prefix}_meta.json")
    metrics_path = Path(f"{prefix}_metrics.txt")
    signals_path = Path(f"{prefix}_signals.csv")
    trades_path = Path(f"{prefix}_trade_log.csv")

    pd.DataFrame(
        {
            "date": r_mk.index.strftime("%Y-%m-%d"),
            "daily_ret_markov": r_mk.values,
            "daily_ret_naive_short": r_naive.values,
            "equity_markov_usd": (cap * (1 + r_mk).cumprod()).values,
            "gross_exposure": win["gross_exposure"].values,
            "cash_weight": win["cash_weight"].values,
        }
    ).to_csv(daily_path, index=False)

    sig_cols = [c for c in win.columns if c.startswith(("weight_", "edge_", "spread_", "decision_"))]
    sig_out = win[sig_cols].copy()
    sig_out.insert(0, "date", sig_out.index.strftime("%Y-%m-%d"))
    sig_out.to_csv(signals_path, index=False)

    if trade_log_df is not None and not trade_log_df.empty:
        t0 = pd.Timestamp(args.start)
        t1 = pd.Timestamp(args.end) if args.end.strip() else None
        ent = pd.to_datetime(trade_log_df["entry_date"])
        trade_log_df = trade_log_df.loc[ent >= t0]
        if t1 is not None:
            trade_log_df = trade_log_df.loc[ent <= t1]
        trade_log_df.to_csv(trades_path, index=False)
        print(f"Wrote {trades_path}  ({len(trade_log_df)} fixed-expiry trades)")

    start_d = str(r_mk.index.min().date()) if len(r_mk) else None
    end_d = str(r_mk.index.max().date()) if len(r_mk) else None

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_markov_vix_spread_backtest.py --start {args.start}"
    )

    meta = {
        "strategy": "markov_vix_calendar_spreads",
        "execution_mode": str(win["execution_mode"].iloc[0]) if "execution_mode" in win.columns else "unknown",
        "spreads": [s[0] for s in available],
        "markov_spread": m_mk,
        "naive_short_all_spreads": m_naive,
        "activity": {
            "short_spread_days": short_days,
            "long_spread_days": long_days,
            "avg_gross_exposure": float(win["gross_exposure"].mean()),
            "n_trades": int(len(trade_log_df)) if trade_log_df is not None else None,
        },
        "start": start_d,
        "end": end_d,
        "capital": cap,
        "artifacts": {
            "daily_csv": str(daily_path),
            "signals_csv": str(signals_path),
            "trade_log_csv": str(trades_path),
        },
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")

    metrics_path.write_text(
        f"# Markov VIX calendar spreads ({start_d} -> {end_d})\n\n"
        f"Command:\n  {cmd}\n\n"
        f"Execution: {meta.get('execution_mode', 'unknown')}\n"
        f"Spreads: {[s[0] for s in available]}\n"
        f"Capital: ${cap:,.0f}\n\n"
        f"--- Markov spread L/S (fixed-expiry) ---\n"
        f"  Total return: {m_mk['total_return_pct']:.2f}%\n"
        f"  CAGR: {m_mk['cagr_pct']:.2f}%\n"
        f"  Sharpe: {m_mk['sharpe']:.3f}\n"
        f"  Max DD: {m_mk['max_drawdown_pct']:.2f}%\n"
        f"  Trades: {len(trade_log_df) if trade_log_df is not None else 'n/a'}\n\n"
        f"--- Naive always-short spreads (constant-tenor benchmark) ---\n"
        f"  Total return: {m_naive['total_return_pct']:.2f}%\n"
        f"  CAGR: {m_naive['cagr_pct']:.2f}%\n"
        f"  Sharpe: {m_naive['sharpe']:.3f}\n"
        f"  Max DD: {m_naive['max_drawdown_pct']:.2f}%\n\n"
        f"Note: prior constant-tenor headline (+5256%) was inflated by rolled-tenor accounting.\n"
    )

    print(f"\n# Markov VIX Calendar Spreads ({start_d} -> {end_d}, {len(r_mk)} days)\n")
    print(f"Spreads: {[s[0] for s in available]}\n")
    print("--- Markov spread L/S (mispricing) ---")
    print(f"  Total return: {m_mk['total_return_pct']:.1f}%")
    print(f"  CAGR: {m_mk['cagr_pct']:.2f}%")
    print(f"  Sharpe: {m_mk['sharpe']:.3f}")
    print(f"  Max DD: {m_mk['max_drawdown_pct']:.2f}%")
    print("\n--- Naive always-short spreads ---")
    print(f"  Total return: {m_naive['total_return_pct']:.1f}%")
    print(f"  CAGR: {m_naive['cagr_pct']:.2f}%")
    print(f"  Sharpe: {m_naive['sharpe']:.3f}")
    print(f"  Max DD: {m_naive['max_drawdown_pct']:.2f}%")
    print(f"\nActivity: long spread days={long_days}, short spread days={short_days}, "
          f"avg gross={win['gross_exposure'].mean():.3f}")

    # Yearly
    print("\n--- Yearly ---")
    yr = pd.DataFrame({"date": r_mk.index, "mk": r_mk.values, "naive": r_naive.values})
    yr["year"] = pd.to_datetime(yr["date"]).dt.year
    for y, g in yr.groupby("year"):
        rm = float((1 + g["mk"]).prod() - 1)
        rn = float((1 + g["naive"]).prod() - 1)
        print(f"  {y}  Markov {rm:+.1%}   naive-short {rn:+.1%}")

    print(f"\nWrote {daily_path}")
    print(f"Wrote {signals_path}")
    print(f"Wrote {meta_path}")

    # Equity curve + trade log + HTML report
    export_script = _REPO / "RenTech" / "strategy_stack" / "export_markov_vix_spread_report.py"
    if export_script.is_file():
        print("\nExporting equity curve and HTML report …", flush=True)
        export_args = [sys.executable, str(export_script), "--prefix", str(prefix), "--capital", str(cap)]
        if trade_log_df is not None and not trade_log_df.empty:
            export_args.append("--engine-trade-log")
        subprocess.run(export_args, cwd=str(_REPO), check=True)


if __name__ == "__main__":
    main()
