#!/usr/bin/env python3
"""
Volatility Edge VIX ETN sleeve (Zarattini / Mele / Aziz, SSRN 5316487).

Implements Strategies 2–4 on short-vol (SVXY/SVIX) and long-vol (VIXY) ETNs:
  * ``evrp`` — short when expected VRP > 0 (20% allocation)
  * ``evrp_boc`` — eVRP + VIX vs VIX3M term structure (Strategy 3)
  * ``evrp_boc_sizing`` — Strategy 4, size = VIX/100
  * ``evrp_boc_sizing_v200`` — Strategy 4 with VIX/200 (conservative)

Signals use daily SPY/VIX/VIX3M closes (paper: 3:45 PM ET + MOC). Rebalance when
target weight drifts > ``rebalance_threshold``; ``tcost_bps`` per traded notional.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_volatility_edge_etn.py \\
      --variant evrp_boc --start 2016-01-04 --end 2026-04-02 --capital 100000

    # Strategy 3 vs 4 (VIX/200) comparison table
    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_volatility_edge_etn.py --compare-s3-s4
"""

from __future__ import annotations

import argparse
import json
import math
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import 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.strategy_stack.data_loader import DataLoader
from RenTech.strategy_stack.spx_regime_state_space.features import _load_vix3m_series
from RenTech.strategy_stack.vrp_backtester import load_spy_vix_from_yfinance

Variant = Literal[
    "evrp",
    "evrp_boc",
    "evrp_boc_sizing",
    "evrp_boc_sizing_v200",
    "passive",
]

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "volatility_edge_etn"
COMPARE_VARIANTS = ("evrp_boc", "evrp_boc_sizing_v200")


@dataclass(frozen=True)
class VolEdgeConfig:
    variant: Variant = "evrp_boc"
    rv_lookback: int = 10
    rebalance_threshold: float = 0.02
    tcost_bps: float = 5.0
    short_ticker: str = "SVXY"
    long_ticker: str = "VIXY"
    short_stitch_ticker: str = "SVIX"
    stitch_date: str = "2022-03-30"


def _daily_returns(ticker: str, index: pd.DatetimeIndex) -> pd.Series:
    loader = DataLoader()
    start = (index.min() - pd.Timedelta(days=30)).strftime("%Y-%m-%d")
    end = (index.max() + pd.Timedelta(days=5)).strftime("%Y-%m-%d")
    df = loader.fetch_daily(ticker, start=start, end=end)
    if df.empty:
        raise RuntimeError(f"No daily data for {ticker}")
    close = df["close"].astype(float)
    close.index = pd.to_datetime(close.index).tz_localize(None)
    return close.reindex(index).ffill().pct_change().fillna(0.0)


def _stitched_short_returns(
    index: pd.DatetimeIndex,
    *,
    primary: str,
    stitch: str,
    stitch_on: pd.Timestamp,
) -> pd.Series:
    """SVXY pre-SVIX inception, SVIX after (short-vol leg)."""
    r_pri = _daily_returns(primary, index)
    if stitch_on <= index.min():
        return r_pri.rename("ret_short_vol")
    r_st = _daily_returns(stitch, index)
    out = r_pri.copy()
    out.loc[index >= stitch_on] = r_st.loc[index >= stitch_on]
    return out.rename("ret_short_vol")


def _erv30(spy_close: pd.Series, i: int, lookback: int) -> float:
    if i < lookback:
        return float("nan")
    window = spy_close.iloc[i - lookback + 1 : i + 1]
    rets = window.pct_change().dropna()
    if len(rets) < lookback - 1:
        return float("nan")
    return float(rets.std(ddof=1) * math.sqrt(252.0) * 100.0)


def target_weights(
    *,
    variant: Variant,
    evrp: float,
    vix: float,
    vix3m: float,
) -> tuple[float, float]:
    """Return (weight_short_vol, weight_long_vol) as fractions of sleeve capital."""
    if not (math.isfinite(evrp) and math.isfinite(vix) and math.isfinite(vix3m)):
        return 0.0, 0.0

    contango = vix < vix3m
    backwardation = vix > vix3m

    if variant == "passive":
        return 0.20, 0.0

    if variant == "evrp":
        return (0.20, 0.0) if evrp > 0 else (0.0, 0.0)

    sizing_div = 100.0
    if variant == "evrp_boc_sizing_v200":
        sizing_div = 200.0

    if variant == "evrp_boc":
        if evrp > 0 and contango:
            return 0.20, 0.0
        if evrp <= 0 and contango:
            return 0.10, 0.0
        if evrp <= 0 and backwardation:
            return 0.0, 0.20
        return 0.0, 0.0

    # dynamic sizing (Strategy 4)
    vix_frac = max(0.0, min(1.0, vix / sizing_div))
    if evrp > 0 and contango:
        return vix_frac, 0.0
    if evrp <= 0 and contango:
        return 0.5 * vix_frac, 0.0
    if evrp <= 0 and backwardation:
        return 0.0, vix_frac
    return 0.0, 0.0


def run_volatility_edge_backtest(
    *,
    idx: pd.DatetimeIndex,
    spy_close: pd.Series,
    vix: pd.Series,
    vix3m: pd.Series,
    ret_short: pd.Series,
    ret_long: pd.Series,
    capital: float,
    cfg: VolEdgeConfig,
) -> tuple[pd.DataFrame, dict]:
    n = len(idx)
    w_short = 0.0
    w_long = 0.0
    rows: list[dict] = []

    for i in range(n):
        dt = idx[i]
        evrp = float(vix.iloc[i]) - _erv30(spy_close, i, cfg.rv_lookback)
        t_short, t_long = target_weights(
            variant=cfg.variant,
            evrp=evrp,
            vix=float(vix.iloc[i]),
            vix3m=float(vix3m.iloc[i]),
        )

        r_s = float(ret_short.iloc[i])
        r_l = float(ret_long.iloc[i])
        port_ret = w_short * r_s + w_long * r_l

        turnover = 0.0
        rebalanced = False
        if (
            abs(w_short - t_short) > cfg.rebalance_threshold
            or abs(w_long - t_long) > cfg.rebalance_threshold
        ):
            turnover = abs(t_short - w_short) + abs(t_long - w_long)
            port_ret -= (cfg.tcost_bps / 10_000.0) * turnover
            w_short, w_long = t_short, t_long
            rebalanced = True
        else:
            # drift weights with ETN returns (cash earns 0)
            gross = w_short * (1.0 + r_s) + w_long * (1.0 + r_l) + max(0.0, 1.0 - w_short - w_long)
            if gross > 1e-12:
                w_short = w_short * (1.0 + r_s) / gross
                w_long = w_long * (1.0 + r_l) / gross

        rows.append(
            {
                "date": dt.strftime("%Y-%m-%d"),
                "daily_ret": port_ret,
                "evrp": evrp,
                "vix": float(vix.iloc[i]),
                "vix3m": float(vix3m.iloc[i]),
                "target_w_short": t_short,
                "target_w_long": t_long,
                "w_short": w_short,
                "w_long": w_long,
                "rebalanced": rebalanced,
                "turnover": turnover,
            }
        )

    daily = pd.DataFrame(rows)
    daily["date"] = pd.to_datetime(daily["date"])
    r = daily["daily_ret"].astype(float)
    cap = float(capital)
    pnl = r * cap
    eq = cap * (1.0 + r).cumprod()
    daily["daily_pnl_usd"] = pnl
    daily["equity_usd"] = eq
    daily["margin_usd"] = eq * (daily["w_short"] + daily["w_long"]) * 0.15

    meta = _metrics(eq, cap, r)
    meta.update(
        {
            "variant": cfg.variant,
            "paper": "Zarattini-Mele-Aziz SSRN 5316487",
            "rebalance_threshold": cfg.rebalance_threshold,
            "tcost_bps": cfg.tcost_bps,
            "short_vol_proxy": f"{cfg.short_ticker} -> {cfg.short_stitch_ticker} from {cfg.stitch_date}",
            "long_vol_proxy": cfg.long_ticker,
            "n_rebalances": int(daily["rebalanced"].sum()),
            "pct_days_short": round(100.0 * (daily["w_short"] > 0.001).mean(), 1),
            "pct_days_long": round(100.0 * (daily["w_long"] > 0.001).mean(), 1),
            "corr_spy": float(
                pd.DataFrame({"v": r, "spy": spy_close.reindex(idx).pct_change().fillna(0.0)})
                .dropna()
                .corr()
                .iloc[0, 1]
            ),
        }
    )
    return daily, meta


def _metrics(eq: pd.Series, cap: float, r: pd.Series) -> dict:
    n = len(r)
    years = n / 252.0
    end = float(eq.iloc[-1])
    tot = end / cap - 1.0
    cagr = (end / cap) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    sd = float(r.std(ddof=1)) if n > 1 else float("nan")
    sharpe = float(r.mean() / sd * math.sqrt(252.0)) if sd > 1e-12 else 0.0
    dd = float((eq / eq.cummax() - 1.0).min())
    return {
        "capital_usd": cap,
        "ending_equity_usd": end,
        "total_return_pct": tot * 100.0,
        "cagr_pct": cagr * 100.0,
        "sharpe": sharpe,
        "max_drawdown_pct": dd * 100.0,
        "vol_ann_pct": sd * math.sqrt(252.0) * 100.0 if math.isfinite(sd) else float("nan"),
        "n_sessions": n,
    }


def load_signal_panel(start: str, end: str, cfg: VolEdgeConfig) -> tuple[pd.DatetimeIndex, pd.DataFrame]:
    yf_start = (pd.Timestamp(start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    yf_end = (pd.Timestamp(end) + pd.Timedelta(days=14)).strftime("%Y-%m-%d")
    spy_vix = load_spy_vix_from_yfinance(yf_start, yf_end)
    spy_vix.index = pd.to_datetime(spy_vix.index).tz_localize(None)
    idx = spy_vix.index[
        (spy_vix.index >= pd.Timestamp(start)) & (spy_vix.index <= pd.Timestamp(end))
    ]
    if len(idx) < 100:
        raise SystemExit(f"Too few sessions ({len(idx)}) for {start} → {end}")

    vix3m = _load_vix3m_series(idx)
    ret_short = _stitched_short_returns(
        idx,
        primary=cfg.short_ticker,
        stitch=cfg.short_stitch_ticker,
        stitch_on=pd.Timestamp(cfg.stitch_date),
    )
    ret_long = _daily_returns(cfg.long_ticker, idx)

    panel = pd.DataFrame(
        {
            "spy_close": spy_vix["close"].reindex(idx).astype(float),
            "vix": spy_vix["vix_close"].reindex(idx).astype(float),
            "vix3m": vix3m.reindex(idx).astype(float),
            "ret_short_vol": ret_short,
            "ret_long_vol": ret_long,
        },
        index=idx,
    )
    panel = panel.dropna(subset=["spy_close", "vix", "vix3m"])
    return pd.DatetimeIndex(panel.index), panel


def run_and_write(
    *,
    start: str,
    end: str,
    capital: float,
    cfg: VolEdgeConfig,
    out_prefix: Path,
) -> dict:
    idx, panel = load_signal_panel(start, end, cfg)
    daily, meta = run_volatility_edge_backtest(
        idx=idx,
        spy_close=panel["spy_close"],
        vix=panel["vix"],
        vix3m=panel["vix3m"],
        ret_short=panel["ret_short_vol"],
        ret_long=panel["ret_long_vol"],
        capital=capital,
        cfg=cfg,
    )
    meta["start"] = str(idx[0].date())
    meta["end"] = str(idx[-1].date())

    prefix = out_prefix.expanduser().resolve()
    if prefix.suffix:
        prefix = prefix.with_suffix("")
    prefix.parent.mkdir(parents=True, exist_ok=True)
    daily_path = Path(f"{prefix}_daily.csv")
    meta_path = Path(f"{prefix}_meta.json")
    daily.to_csv(daily_path, index=False)
    meta_path.write_text(json.dumps(meta, indent=2) + "\n", encoding="utf-8")

    print(
        f"  {cfg.variant}: return {meta['total_return_pct']:+.1f}%  "
        f"CAGR {meta['cagr_pct']:+.1f}%  Sharpe {meta['sharpe']:.2f}  "
        f"maxDD {meta['max_drawdown_pct']:.1f}%  ρ(SPY) {meta['corr_spy']:.2f}",
        flush=True,
    )
    print(f"    → {daily_path}", flush=True)
    return meta


def compare_s3_s4(*, start: str, end: str, capital: float) -> None:
    print(f"\n=== Volatility Edge standalone · {start} → {end} ===", flush=True)
    results: list[dict] = []
    for variant in COMPARE_VARIANTS:
        cfg = VolEdgeConfig(variant=variant)  # type: ignore[arg-type]
        out = DEFAULT_OUT_PREFIX.parent / f"volatility_edge_etn_{variant}"
        meta = run_and_write(start=start, end=end, capital=capital, cfg=cfg, out_prefix=out)
        results.append(meta)

    # quick combine with stock-only daily (if present)
    stock_path = LOGS / (
        "stock_only_tactical70_plus_stock_only_plus_equity_dip_plus_sector_momentum_"
        "plus_tactical_aw_plus_tsmom_plus_fund_plus_nav_q_mtm_daily.csv"
    )
    if stock_path.is_file():
        print("\n=== Stock-only + Vol Edge blends (static weight add) ===", flush=True)
        sdf = pd.read_csv(stock_path, parse_dates=["date"]).set_index("date").sort_index()
        sret = sdf["equity_mtm_usd"].astype(float).pct_change().fillna(0.0)
        for variant in COMPARE_VARIANTS:
            vdf = pd.read_csv(
                LOGS / f"volatility_edge_etn_{variant}_daily.csv", parse_dates=["date"]
            ).set_index("date")
            vret = vdf["daily_ret"].astype(float)
            for vw in (0.10, 0.125, 0.15):
                w = 1.0 - vw
                blend = w * sret.reindex(sret.index.union(vret.index)).fillna(0.0) + vw * vret.reindex(
                    sret.index.union(vret.index)
                ).fillna(0.0)
                eq = capital * (1.0 + blend).cumprod()
                r = eq.pct_change().fillna(0.0)
                sh = float(r.mean() / r.std(ddof=1) * math.sqrt(252)) if r.std() > 0 else 0.0
                dd = float((eq / eq.cummax() - 1.0).min())
                tot = float(eq.iloc[-1] / capital - 1.0) * 100.0
                print(
                    f"  {variant} @ {vw:.0%} vol: return {tot:+.1f}%  Sharpe {sh:.2f}  maxDD {dd*100:.1f}%",
                    flush=True,
                )
        results.append({"note": "blend uses existing stock_only_tactical70 fund daily"})

    out_json = LOGS / "volatility_edge_etn_compare_s3_s4.json"
    out_json.write_text(json.dumps(results, indent=2) + "\n", encoding="utf-8")
    print(f"\nWrote {out_json}", flush=True)


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-06-18")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument(
        "--variant",
        default="evrp_boc",
        choices=[
            "passive",
            "evrp",
            "evrp_boc",
            "evrp_boc_sizing",
            "evrp_boc_sizing_v200",
        ],
    )
    ap.add_argument("--rebalance-threshold", type=float, default=0.02)
    ap.add_argument("--tcost-bps", type=float, default=5.0)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument(
        "--compare-s3-s4",
        action="store_true",
        help="Run Strategy 3 vs Strategy 4 (VIX/200) and write comparison JSON",
    )
    ap.add_argument(
        "--all-variants",
        action="store_true",
        help="Run evrp, evrp_boc, evrp_boc_sizing, evrp_boc_sizing_v200",
    )
    args = ap.parse_args()

    if args.compare_s3_s4:
        compare_s3_s4(start=str(args.start), end=str(args.end), capital=float(args.capital))
        return

    variants: list[Variant]
    if args.all_variants:
        variants = ["evrp", "evrp_boc", "evrp_boc_sizing", "evrp_boc_sizing_v200"]
    else:
        variants = [args.variant]  # type: ignore[list-item]

    for variant in variants:
        cfg = VolEdgeConfig(
            variant=variant,
            rebalance_threshold=float(args.rebalance_threshold),
            tcost_bps=float(args.tcost_bps),
        )
        prefix = args.out_prefix
        if len(variants) > 1 or variant != "evrp_boc":
            prefix = args.out_prefix.parent / f"{args.out_prefix.name}_{variant}"
        run_and_write(
            start=str(args.start),
            end=str(args.end),
            capital=float(args.capital),
            cfg=cfg,
            out_prefix=prefix,
        )


if __name__ == "__main__":
    main()
