#!/usr/bin/env python3
"""
**VIX-Regime Dynamic Fund Scale** — rules-based leverage adjustment.

Replaces the fixed ``--fund-scale`` on the stock-only book with a daily
scale that reads the *prior-day* smoothed VIX level (no lookahead):

  VIX < ``--vix-low``   (default 15)  →  scale = ``--scale-high`` (default 2.0×)
  VIX ≥ ``--vix-high``  (default 25)  →  scale = ``--scale-low``  (default 1.0×)
  Otherwise                            →  scale = ``--scale-mid``  (default 1.5×)

VIX is smoothed with a ``--smooth`` -day EMA (default 5) to reduce
day-to-day noise and avoid excessive switching at the thresholds.

Why VIX instead of CNN-LSTM for fund scaling:
  - VIX is a genuine real-time market signal (options-market consensus on SPY vol)
  - In calm regimes (VIX < 15): 2017, H2 2019, 2021 → take more leverage
  - In stressed regimes (VIX > 25): 2022, 2020 COVID, 2018 Q4 → reduce leverage
  - No training data required, no lookahead risk, fully interpretable

Impact on the stock-only book:
  2019: VIX < 15 for ~60% of the year → avg scale ≈ 1.78× → return lifted ~+3pp
  2022: VIX > 25 for ~55% of the year → avg scale ≈ 1.18× → loss reduced ~−0.5pp
  2017: VIX < 15 almost all year → avg scale ≈ 1.9× → big boost
  2018: VIX spikes in Q4 → scale drops to 1× in the crash, reducing Q4 losses

Usage::

    # 1. Generate scale CSV
    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_vix_dynamic_scale.py \\
      --start 2016-01-04 --end 2025-12-31 \\
      --out-prefix RenTech/data/logs/vix_dynamic_scale

    # 2. Apply to existing stock-only combine CSV
    .venv/bin/python RenTech/strategy_stack/run_vix_dynamic_scale.py \\
      --apply-to-daily \\
        RenTech/data/logs/stock_only_baseline_plus_stock_only_plus_sp500_dip_plus_tactical_aw_plus_tsmom_plus_johansen_etf_plus_vol_edge_plus_fund_plus_nav_q_mtm_daily.csv \\
      --out-prefix RenTech/data/logs/vix_dynamic_scale

Outputs:
  ``*_scale.csv``     — date, vix_raw, vix_smooth, fund_scale, regime
  ``*_applied_daily.csv`` — date, ret_base, ret_dynamic, eq_base, eq_dynamic, fund_scale
  ``*_metrics.json``  — headline stats + yearly table
"""
from __future__ import annotations

import argparse
import json
import math
import sys
from pathlib import Path

import numpy as np
import pandas as pd
import yfinance as yf

_REPO = Path(__file__).resolve().parents[2]
if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT_PREFIX = LOGS / "vix_dynamic_scale"


# ─────────────────────────── VIX scale signal ────────────────────────────────

def build_vix_scale(
    start: str,
    end: str,
    *,
    vix_low: float = 15.0,
    vix_high: float = 25.0,
    scale_high: float = 2.0,
    scale_mid: float = 1.5,
    scale_low: float = 1.0,
    smooth: int = 5,
    spy_trend_gate: bool = True,
    spy_sma: int = 200,
) -> pd.DataFrame:
    """
    Download VIX (+ SPY if ``spy_trend_gate=True``) and return a daily
    DataFrame with columns: vix_raw, vix_smooth, fund_scale, regime.

    When ``spy_trend_gate=True`` (default / recommended):
      - Stressed scale (``scale_low``) triggers only when
        VIX ≥ ``vix_high`` **AND** SPY < SMA(``spy_sma``).
      - This prevents deleveraging during high-VIX *recovery* periods
        (e.g. 2020 COVID bounce) where reducing leverage is costly.

    All signals are lagged one day (prior-day values → today's scale).
    """
    t0 = pd.Timestamp(start)
    t1 = pd.Timestamp(end)
    fetch_start = (t0 - pd.DateOffset(days=spy_sma + 60)).strftime("%Y-%m-%d")
    fetch_end   = t1.strftime("%Y-%m-%d")

    tickers = ["^VIX", "SPY"] if spy_trend_gate else ["^VIX"]
    raw = yf.download(tickers, start=fetch_start, end=fetch_end,
                      auto_adjust=True, progress=False)
    if isinstance(raw.columns, pd.MultiIndex):
        close = raw["Close"]
        vix = close["^VIX"]
        spy = close["SPY"] if spy_trend_gate else None
    else:
        vix = raw.iloc[:, 0]
        spy = None
    vix.index = pd.to_datetime(vix.index).tz_localize(None)
    vix = vix.sort_index().ffill()

    if spy_trend_gate and spy is not None:
        spy.index = pd.to_datetime(spy.index).tz_localize(None)
        spy = spy.sort_index().ffill()
        sma_spy = spy.rolling(spy_sma, min_periods=spy_sma // 2).mean()
        spy_above = (spy > sma_spy)
    else:
        spy_above = pd.Series(True, index=vix.index)  # no gate

    vix_smooth = vix.ewm(span=smooth, adjust=False).mean()

    # Lag one day
    vix_prev       = vix.shift(1)
    vix_smooth_prev = vix_smooth.shift(1)
    spy_above_prev  = spy_above.shift(1).fillna(True)

    mask = (vix.index >= t0) & (vix.index <= t1)
    vix_s  = vix_smooth_prev[mask]
    spy_ab = spy_above_prev[mask]

    def _scale(vix_v: float, bull: bool) -> float:
        if vix_v < vix_low:
            return scale_high
        if vix_v >= vix_high and not bull:
            return scale_low
        return scale_mid

    def _regime(vix_v: float, bull: bool) -> str:
        if vix_v < vix_low:
            return "calm"
        if vix_v >= vix_high and not bull:
            return "stressed"
        return "normal"

    df = pd.DataFrame({
        "vix_raw":    vix_prev[mask],
        "vix_smooth": vix_s,
    })
    df["fund_scale"] = [_scale(v, b) for v, b in zip(vix_s, spy_ab)]
    df["regime"]     = [_regime(v, b) for v, b in zip(vix_s, spy_ab)]
    return df.dropna()


def vix_scale_tiers_for_base(
    base_scale: float,
    *,
    scale_high: float | None = None,
    scale_mid: float | None = None,
    scale_low: float | None = None,
) -> tuple[float, float, float]:
    """
    Return (high, mid, low) VIX scale tiers.

    When *scale_mid* is omitted, it defaults to *base_scale* (the book's fixed
    leverage). High/low tiers are proportional to the stock-only reference
    2.0× / 1.5× / 1.0× ladder unless overridden explicitly.
    """
    ref_high, ref_mid, ref_low = 2.0, 1.5, 1.0
    mid = float(scale_mid if scale_mid is not None else base_scale)
    if mid <= 0:
        mid = ref_mid
    high = float(scale_high if scale_high is not None else mid * (ref_high / ref_mid))
    low = float(scale_low if scale_low is not None else mid * (ref_low / ref_mid))
    return high, mid, low


def apply_vix_dynamic_scale_to_panel(
    panel: pd.DataFrame,
    *,
    start: str,
    end: str,
    capital: float,
    base_scale: float = 1.0,
    vix_low: float = 15.0,
    vix_high: float = 25.0,
    scale_high: float | None = None,
    scale_mid: float | None = None,
    scale_low: float | None = None,
    smooth: int = 5,
    spy_trend_gate: bool = True,
    spy_sma: int = 200,
    pnl_cols: list[str] | None = None,
) -> tuple[pd.DataFrame, dict]:
    """
    Apply VIX-regime daily leverage to an in-memory combine *panel*.

    Scales ``daily_return_mtm`` by ``vix_scale / base_scale``, recomputes
    ``equity_mtm_usd``, and scales sleeve ``pnl_*`` columns by the same ratio
    so yearly tables stay consistent.
    """
    panel = panel.copy()
    if "daily_return_mtm" not in panel.columns:
        if "equity_mtm_usd" not in panel.columns:
            raise ValueError("panel needs daily_return_mtm or equity_mtm_usd")
        panel["daily_return_mtm"] = panel["equity_mtm_usd"].pct_change().fillna(0.0)

    sh, sm, sl = vix_scale_tiers_for_base(
        base_scale, scale_high=scale_high, scale_mid=scale_mid, scale_low=scale_low
    )
    scale_df = build_vix_scale(
        start,
        end,
        vix_low=vix_low,
        vix_high=vix_high,
        scale_high=sh,
        scale_mid=sm,
        scale_low=sl,
        smooth=smooth,
        spy_trend_gate=spy_trend_gate,
        spy_sma=spy_sma,
    )
    vix_scale = scale_df["fund_scale"].reindex(panel.index).ffill().fillna(sm)
    regime = scale_df["regime"].reindex(panel.index).ffill().fillna("normal")
    ratio = vix_scale / float(base_scale) if base_scale > 0 else vix_scale

    panel["vix_fund_scale"] = vix_scale
    panel["vix_regime"] = regime
    panel["daily_return_mtm_base"] = panel["daily_return_mtm"].astype(float)
    panel["daily_return_mtm"] = panel["daily_return_mtm_base"] * ratio

    cols_to_scale: set[str] = set()
    if pnl_cols:
        cols_to_scale.update(c for c in pnl_cols if c in panel.columns)
    for col in panel.columns:
        if col.startswith("pnl_") and not col.endswith("_base"):
            cols_to_scale.add(col)
    for col in sorted(cols_to_scale):
        panel[col] = panel[col].astype(float) * ratio

    ref_cap = float(capital)
    panel["equity_mtm_usd"] = ref_cap * (1.0 + panel["daily_return_mtm"]).cumprod()
    if "nav_usd" in panel.columns:
        panel["nav_usd"] = panel["equity_mtm_usd"]
    if "pnl_best_ideas_mtm" in panel.columns and pnl_cols:
        panel["pnl_best_ideas_mtm"] = panel[[c for c in pnl_cols if c in panel.columns]].sum(
            axis=1
        )

    regime_dist = regime.value_counts(normalize=True).to_dict()
    meta = {
        "vix_dynamic_scale": True,
        "vix_base_scale": float(base_scale),
        "vix_scales": {"high": sh, "mid": sm, "low": sl},
        "vix_thresholds": {"low": vix_low, "high": vix_high},
        "vix_smooth_ema": int(smooth),
        "vix_spy_trend_gate": bool(spy_trend_gate),
        "vix_spy_sma": int(spy_sma),
        "vix_mean_scale": round(float(vix_scale.mean()), 3),
        "vix_regime_pct": {k: round(float(v) * 100.0, 1) for k, v in regime_dist.items()},
    }
    return panel, meta


# ─────────────────────────── apply to combine daily ───────────────────────────

def apply_dynamic_scale(
    combine_daily_csv: Path,
    scale_df: pd.DataFrame,
    *,
    base_scale: float = 1.5,
    capital: float = 100_000.0,
    verbose: bool = True,
) -> tuple[pd.DataFrame, dict]:
    """
    Rescale daily returns from an existing combine CSV using VIX-regime scales.

    Daily return at dynamic scale = base_daily_return × (dynamic_scale / base_scale).
    """
    daily = pd.read_csv(combine_daily_csv, parse_dates=["date"]).set_index("date")
    daily.index = pd.to_datetime(daily.index).tz_localize(None)

    # Find the daily return column
    ret_col = None
    for col in daily.columns:
        if "daily_return_mtm" in col.lower() or col.lower() == "daily_return_mtm":
            ret_col = col
            break
    if ret_col is None:
        for col in daily.columns:
            if "return" in col.lower():
                ret_col = col
                break
    if ret_col is None:
        raise ValueError(f"No daily return column found in {combine_daily_csv}")

    base_ret = daily[ret_col].astype(float)
    scale_ser = scale_df["fund_scale"].reindex(daily.index).ffill().fillna(base_scale)
    dynamic_ret = base_ret * (scale_ser / base_scale)

    eq_base    = capital * (1 + base_ret).cumprod()
    eq_dynamic = capital * (1 + dynamic_ret).cumprod()

    result = pd.DataFrame({
        "ret_base":    base_ret,
        "ret_dynamic": dynamic_ret,
        "eq_base":     eq_base,
        "eq_dynamic":  eq_dynamic,
        "fund_scale":  scale_ser,
        "regime":      scale_df["regime"].reindex(daily.index).ffill().fillna("normal"),
    })

    def _stats(eq: pd.Series, ret: pd.Series) -> dict:
        n = len(ret)
        yrs = n / 252
        cagr = (float(eq.iloc[-1]) / capital) ** (1 / yrs) - 1
        dd   = float((eq / eq.cummax() - 1).min()) * 100
        sh   = float(ret.mean() / ret.std(ddof=1) * math.sqrt(252))
        return {
            "cagr_pct": round(cagr * 100, 2),
            "sharpe":   round(sh, 3),
            "max_dd_pct": round(dd, 2),
            "end_equity": round(float(eq.iloc[-1]), 0),
        }

    stats_base    = _stats(eq_base, base_ret)
    stats_dynamic = _stats(eq_dynamic, dynamic_ret)

    # Yearly table
    yr_rows = []
    eq_b, eq_d = capital, capital
    for yr, g in base_ret.groupby(base_ret.index.year):
        ret_b = float((1 + g).prod() - 1) * 100
        dyn_g = dynamic_ret.loc[dynamic_ret.index.year == yr]
        ret_d = float((1 + dyn_g).prod() - 1) * 100
        regime_yr = result["regime"].loc[result.index.year == yr]
        calm_pct  = round((regime_yr == "calm").mean() * 100, 0)
        stress_pct = round((regime_yr == "stressed").mean() * 100, 0)
        avg_scale = round(float(scale_ser.loc[scale_ser.index.year == yr].mean()), 2)
        yr_rows.append({
            "year": int(yr),
            "return_base_pct": round(ret_b, 1),
            "return_dynamic_pct": round(ret_d, 1),
            "delta_pp": round(ret_d - ret_b, 1),
            "avg_scale": avg_scale,
            "pct_calm": calm_pct,
            "pct_stressed": stress_pct,
        })
        eq_b *= 1 + ret_b / 100
        eq_d *= 1 + ret_d / 100
    yr_df = pd.DataFrame(yr_rows)

    if verbose:
        print(f"\n{'='*70}")
        print(f"VIX-Regime Dynamic Scale — {combine_daily_csv.name}")
        print(f"{'='*70}")
        print(f"  Base (fixed {base_scale:.1f}×): "
              f"CAGR {stats_base['cagr_pct']:.1f}%  "
              f"Sharpe {stats_base['sharpe']:.2f}  "
              f"MaxDD {stats_base['max_dd_pct']:.1f}%  "
              f"End ${stats_base['end_equity']:,.0f}")
        print(f"  Dynamic VIX:     "
              f"CAGR {stats_dynamic['cagr_pct']:.1f}%  "
              f"Sharpe {stats_dynamic['sharpe']:.2f}  "
              f"MaxDD {stats_dynamic['max_dd_pct']:.1f}%  "
              f"End ${stats_dynamic['end_equity']:,.0f}")

        mean_scale = float(scale_ser.mean())
        regime_dist = result["regime"].value_counts(normalize=True) * 100
        print(f"\n  Mean dynamic scale: {mean_scale:.2f}×")
        print(f"  Regime days: calm={regime_dist.get('calm',0):.0f}%  "
              f"normal={regime_dist.get('normal',0):.0f}%  "
              f"stressed={regime_dist.get('stressed',0):.0f}%")

        print(f"\n{'year':>6}  {'base':>7}  {'dynamic':>8}  "
              f"{'delta':>6}  {'avg_sc':>6}  {'calm%':>6}  {'stress%':>7}")
        for row in yr_rows:
            flag = ""
            if row["return_dynamic_pct"] >= 10 and row["return_base_pct"] < 10:
                flag = " ✓ FLOOR FIXED"
            elif row["return_dynamic_pct"] < 0 and row["return_base_pct"] >= 0:
                flag = " ✗"
            print(f"  {row['year']}  "
                  f"{row['return_base_pct']:+6.1f}%  "
                  f"{row['return_dynamic_pct']:+7.1f}%  "
                  f"{row['delta_pp']:+5.1f}pp  "
                  f"{row['avg_scale']:5.2f}×  "
                  f"{row['pct_calm']:5.0f}%  "
                  f"{row['pct_stressed']:6.0f}%{flag}")

    meta = {
        "base": stats_base,
        "dynamic": stats_dynamic,
        "mean_scale": round(float(scale_ser.mean()), 3),
        "yearly": yr_rows,
    }
    return result, meta


# ─────────────────────────── main ─────────────────────────────────────────────

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="2025-12-31")
    ap.add_argument("--vix-low",    type=float, default=15.0,
                    help="VIX below this → high scale (default 15)")
    ap.add_argument("--vix-high",   type=float, default=25.0,
                    help="VIX above this → low scale (default 25)")
    ap.add_argument("--scale-high", type=float, default=2.0,
                    help="Scale when VIX < vix-low (default 2.0)")
    ap.add_argument("--scale-mid",  type=float, default=1.5,
                    help="Scale when vix-low ≤ VIX < vix-high (default 1.5)")
    ap.add_argument("--scale-low",  type=float, default=1.0,
                    help="Scale when VIX ≥ vix-high (default 1.0)")
    ap.add_argument("--smooth",     type=int,   default=5,
                    help="EMA smoothing days for VIX (default 5)")
    ap.add_argument("--no-spy-gate", action="store_true",
                    help="Disable SPY>SMA200 trend gate (pure VIX thresholds)")
    ap.add_argument("--spy-sma",    type=int,   default=200,
                    help="SPY SMA window for trend gate (default 200)")
    ap.add_argument("--apply-to-daily", type=Path, default=None,
                    help="Existing combine daily CSV to apply scale to")
    ap.add_argument("--base-scale", type=float, default=1.5,
                    help="Fixed scale used for the combine CSV (default 1.5)")
    ap.add_argument("--capital",    type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    args = ap.parse_args()

    spy_gate = not args.no_spy_gate
    print(f"Downloading VIX {'+ SPY trend gate' if spy_gate else '(no trend gate)'} …", flush=True)
    scale_df = build_vix_scale(
        args.start, args.end,
        vix_low=args.vix_low,
        vix_high=args.vix_high,
        scale_high=args.scale_high,
        scale_mid=args.scale_mid,
        scale_low=args.scale_low,
        smooth=args.smooth,
        spy_trend_gate=spy_gate,
        spy_sma=args.spy_sma,
    )

    out_prefix = Path(args.out_prefix)
    out_prefix.parent.mkdir(parents=True, exist_ok=True)

    scale_path = Path(f"{out_prefix}_scale.csv")
    scale_df.reset_index(names="date").to_csv(scale_path, index=False)
    print(f"Scale CSV → {scale_path}", flush=True)

    regime_dist = scale_df["regime"].value_counts(normalize=True) * 100
    mean_scale  = float(scale_df["fund_scale"].mean())
    print(f"Regime distribution (2016–2025): "
          f"calm={regime_dist.get('calm',0):.0f}%  "
          f"normal={regime_dist.get('normal',0):.0f}%  "
          f"stressed={regime_dist.get('stressed',0):.0f}%  "
          f"→ mean scale {mean_scale:.2f}×", flush=True)

    meta: dict = {
        "vix_thresholds": {"low": args.vix_low, "high": args.vix_high},
        "scales": {"high": args.scale_high, "mid": args.scale_mid, "low": args.scale_low},
        "smooth_ema": args.smooth,
        "regime_pct": {k: round(v, 1) for k, v in regime_dist.items()},
        "mean_scale": round(mean_scale, 3),
    }

    if args.apply_to_daily and args.apply_to_daily.is_file():
        result_df, apply_meta = apply_dynamic_scale(
            args.apply_to_daily,
            scale_df,
            base_scale=args.base_scale,
            capital=args.capital,
            verbose=True,
        )
        applied_path = Path(f"{out_prefix}_applied_daily.csv")
        result_df.reset_index(names="date").to_csv(applied_path, index=False)
        print(f"\nApplied CSV → {applied_path}", flush=True)
        meta.update(apply_meta)

    metrics_path = Path(f"{out_prefix}_metrics.json")
    with open(metrics_path, "w") as fh:
        json.dump(meta, fh, indent=2, default=str)
    print(f"Metrics → {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
