#!/usr/bin/env python3
"""
Explore **VXX forward returns** vs signal combinations:

* **VXX vs moving average** — ``vxx_vs_ma`` = VXX / SMA(window) − 1
* **VX1 vs moving average** — ``vx1_elev_vs_ma`` (front-month futures level)
* **Roll cost** — ``roll_cost_m1_m2`` = VX2/VX1 − 1; optional ``roll_cost_vx3_vx1``

For each horizon *N* (trading days), computes forward simple return and summarizes
by quantile bins, threshold grids, and custom filter combos.

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/analyze_vxx_forward_returns.py \\
        --start 2016-01-01 --end 2026-12-31 \\
        --forward-days 5 10 20 40 --ma-window 60
    PYTHONUNBUFFERED=1 .venv/bin/python RenTech/strategy_stack/analyze_vxx_forward_returns.py \\
        --start 2020-01-01 --end 2025-12-31 \\
        --forward-days 10 20 --grid-roll 0.05 0.07 0.09 0.11 \\
        --grid-elev 0.0 0.02 0.05 0.08 --out-csv RenTech/data/logs/vxx_fwd_return_grid.csv
"""

from __future__ import annotations

import argparse
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))

DATA_DIR = _REPO / "RenTech" / "data"
LOGS = DATA_DIR / "logs"
VIX_PANEL = DATA_DIR / "vix_futures_cboe.parquet"


def load_vxx_close(start: str, end: str) -> pd.Series:
    import yfinance as yf

    raw = yf.download("VXX", start=start, end=end, progress=False, auto_adjust=True)
    if raw.empty:
        raise RuntimeError("No VXX data from yfinance")
    close = raw["Close"]
    if isinstance(close, pd.DataFrame):
        close = close.squeeze()
    close.index = pd.to_datetime(close.index).tz_localize(None).normalize()
    close = close.sort_index()
    close.name = "vxx_close"
    return close.loc[start:end]


def load_vix_signals(ma_window: int) -> pd.DataFrame:
    if not VIX_PANEL.is_file():
        raise FileNotFoundError(f"Missing {VIX_PANEL}; run download_cboe_vix_futures.py")
    ct = pd.read_parquet(VIX_PANEL)
    ct.index = pd.to_datetime(ct.index).normalize()
    s1 = ct["vx1_settle"].fillna(ct.get("vx1_close"))
    s2 = ct["vx2_settle"].fillna(ct.get("vx2_close"))
    s3 = ct.get("vx3_settle", pd.Series(dtype=float)).fillna(ct.get("vx3_close"))
    out = pd.DataFrame(index=ct.index)
    out["vx1"] = s1
    out["vix_spot"] = ct.get("vix_spot")
    out["roll_cost_m1_m2"] = (s2 / s1) - 1.0
    if s3 is not None and s3.notna().any():
        out["roll_cost_vx3_vx1"] = (s3 / s1) - 1.0
    min_p = max(ma_window // 2, 20)
    out["vx1_sma"] = out["vx1"].rolling(ma_window, min_periods=min_p).mean()
    out["vx1_elev_vs_ma"] = (out["vx1"] / out["vx1_sma"]) - 1.0
    out["vix_elev_vs_ma"] = (out["vix_spot"] / out["vix_spot"].rolling(ma_window, min_periods=min_p).mean()) - 1.0
    return out


def build_panel(
    start: str,
    end: str,
    *,
    ma_window: int,
    forward_days: list[int],
) -> pd.DataFrame:
    vxx = load_vxx_close(start, end)
    vix = load_vix_signals(ma_window)
    df = pd.DataFrame({"vxx_close": vxx})
    df = df.join(vix, how="left")
    min_p = max(ma_window // 2, 20)
    df["vxx_sma"] = df["vxx_close"].rolling(ma_window, min_periods=min_p).mean()
    df["vxx_vs_ma"] = (df["vxx_close"] / df["vxx_sma"]) - 1.0

    for n in forward_days:
        df[f"fwd_ret_{n}d"] = df["vxx_close"].shift(-n) / df["vxx_close"] - 1.0
        df[f"fwd_log_ret_{n}d"] = np.log(df["vxx_close"].shift(-n) / df["vxx_close"])

    return df.loc[start:end].copy()


def _assign_quantile_bins(s: pd.Series, n_bins: int, labels: list[str] | None = None) -> pd.Series:
    valid = s.dropna()
    if len(valid) < n_bins * 5:
        return pd.Series(pd.NA, index=s.index, dtype="object")
    try:
        return pd.qcut(s, q=n_bins, duplicates="drop", labels=labels)
    except ValueError:
        return pd.qcut(s.rank(method="first"), q=n_bins, labels=labels)


def summarize_quantile_bins(
    df: pd.DataFrame,
    forward_n: int,
    *,
    n_bins: int = 5,
) -> pd.DataFrame:
    """Univariate: each signal quantile → forward return stats."""
    col_ret = f"fwd_ret_{forward_n}d"
    rows = []
    signals = [
        ("vxx_vs_ma", "VXX vs MA"),
        ("vx1_elev_vs_ma", "VX1 vs MA"),
        ("roll_cost_m1_m2", "Roll M1→M2"),
        ("roll_cost_vx3_vx1", "Roll VX3/VX1"),
        ("vix_elev_vs_ma", "VIX spot vs MA"),
    ]
    for sig_col, sig_label in signals:
        if sig_col not in df.columns:
            continue
        sub = df[[sig_col, col_ret]].dropna()
        if sub.empty:
            continue
        labels = [f"Q{i+1}" for i in range(min(n_bins, 5))]
        sub = sub.copy()
        sub["bin"] = _assign_quantile_bins(sub[sig_col], n_bins, labels[:n_bins])
        sub = sub.dropna(subset=["bin"])
        for b, g in sub.groupby("bin", observed=True):
            r = g[col_ret]
            rows.append({
                "forward_days": forward_n,
                "signal": sig_label,
                "bin": str(b),
                "n": len(g),
                "mean_fwd_pct": round(float(r.mean()) * 100, 3),
                "median_fwd_pct": round(float(r.median()) * 100, 3),
                "std_fwd_pct": round(float(r.std()) * 100, 3),
                "hit_neg_pct": round(float((r < 0).mean()) * 100, 1),
                "hit_pos_pct": round(float((r > 0).mean()) * 100, 1),
                "sig_lo": round(float(g[sig_col].min()) * 100, 2),
                "sig_hi": round(float(g[sig_col].max()) * 100, 2),
            })
    return pd.DataFrame(rows)


def summarize_2d_grid(
    df: pd.DataFrame,
    forward_n: int,
    *,
    roll_col: str = "roll_cost_m1_m2",
    elev_col: str = "vx1_elev_vs_ma",
    roll_edges: list[float],
    elev_edges: list[float],
) -> pd.DataFrame:
    """2D threshold grid: roll × elevation → forward return."""
    col_ret = f"fwd_ret_{forward_n}d"
    sub = df[[roll_col, elev_col, col_ret]].dropna()
    if sub.empty:
        return pd.DataFrame()

    roll_edges = sorted(roll_edges)
    elev_edges = sorted(elev_edges)
    rows = []
    roll_breaks = [-np.inf] + roll_edges + [np.inf]
    elev_breaks = [-np.inf] + elev_edges + [np.inf]

    for i in range(len(roll_breaks) - 1):
        for j in range(len(elev_breaks) - 1):
            rlo, rhi = roll_breaks[i], roll_breaks[i + 1]
            elo, ehi = elev_breaks[j], elev_breaks[j + 1]
            mask = (
                (sub[roll_col] >= rlo)
                & (sub[roll_col] < rhi)
                & (sub[elev_col] >= elo)
                & (sub[elev_col] < ehi)
            )
            g = sub.loc[mask, col_ret]
            if len(g) < 3:
                continue
            rows.append({
                "forward_days": forward_n,
                "roll_lo": rlo if np.isfinite(rlo) else None,
                "roll_hi": rhi if np.isfinite(rhi) else None,
                "elev_lo": elo if np.isfinite(elo) else None,
                "elev_hi": ehi if np.isfinite(ehi) else None,
                "roll_label": _fmt_band(rlo, rhi, pct=True),
                "elev_label": _fmt_band(elo, ehi, pct=True),
                "n": len(g),
                "mean_fwd_pct": round(float(g.mean()) * 100, 3),
                "median_fwd_pct": round(float(g.median()) * 100, 3),
                "hit_neg_pct": round(float((g < 0).mean()) * 100, 1),
                "sharpe_fwd": round(float(g.mean() / g.std()) * np.sqrt(252 / forward_n), 3)
                if g.std() > 0
                else None,
            })
    return pd.DataFrame(rows)


def summarize_combo_filters(
    df: pd.DataFrame,
    forward_n: int,
    combos: list[dict],
) -> pd.DataFrame:
    """Named boolean filters → forward return stats."""
    col_ret = f"fwd_ret_{forward_n}d"
    rows = []
    for c in combos:
        mask = pd.Series(True, index=df.index)
        for col, op, val in c.get("filters", []):
            if col not in df.columns:
                mask &= False
                continue
            s = df[col]
            if op == ">=":
                mask &= s >= val
            elif op == "<=":
                mask &= s <= val
            elif op == ">":
                mask &= s > val
            elif op == "<":
                mask &= s < val
        g = df.loc[mask, col_ret].dropna()
        rows.append({
            "forward_days": forward_n,
            "name": c["name"],
            "n": len(g),
            "mean_fwd_pct": round(float(g.mean()) * 100, 3) if len(g) else None,
            "median_fwd_pct": round(float(g.median()) * 100, 3) if len(g) else None,
            "hit_neg_pct": round(float((g < 0).mean()) * 100, 1) if len(g) else None,
            "std_fwd_pct": round(float(g.std()) * 100, 3) if len(g) > 1 else None,
        })
    return pd.DataFrame(rows)


def _fmt_band(lo: float, hi: float, *, pct: bool = False) -> str:
    def f(x: float) -> str:
        if not np.isfinite(x):
            return "∞"
        return f"{x*100:.1f}%" if pct else f"{x:.3f}"

    return f"[{f(lo)}, {f(hi)})"


def default_combos() -> list[dict]:
    return [
        {
            "name": "high_roll_vx1_elev (long put thesis)",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.07),
                ("vx1_elev_vs_ma", ">=", 0.0),
            ],
        },
        {
            "name": "user_9pct_roll_vx1_elev",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.09),
                ("vx1_elev_vs_ma", ">=", 0.0),
            ],
        },
        {
            "name": "high_roll_vxx_above_ma",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.07),
                ("vxx_vs_ma", ">=", 0.0),
            ],
        },
        {
            "name": "steep_roll_vxx_below_ma",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.07),
                ("vxx_vs_ma", "<", 0.0),
            ],
        },
        {
            "name": "low_roll_backwardation",
            "filters": [("roll_cost_m1_m2", "<", 0.0)],
        },
        {
            "name": "vxx_well_above_ma_no_roll",
            "filters": [
                ("vxx_vs_ma", ">=", 0.05),
                ("roll_cost_m1_m2", "<", 0.03),
            ],
        },
        {
            "name": "sweet_spot_roll7_vx1_mild_elev",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.07),
                ("roll_cost_m1_m2", "<", 0.12),
                ("vx1_elev_vs_ma", ">=", 0.0),
                ("vx1_elev_vs_ma", "<=", 0.05),
            ],
        },
        {
            "name": "sweet_spot_roll7_vx1_mild_elev_vxx_cap",
            "filters": [
                ("roll_cost_m1_m2", ">=", 0.07),
                ("roll_cost_m1_m2", "<", 0.12),
                ("vx1_elev_vs_ma", ">=", 0.0),
                ("vx1_elev_vs_ma", "<=", 0.05),
                ("vxx_vs_ma", "<=", 0.05),
            ],
        },
    ]


def print_section(title: str) -> None:
    print("\n" + "=" * 90)
    print(title)
    print("=" * 90)


def main() -> None:
    ap = argparse.ArgumentParser(description="VXX forward returns vs MA + roll signals")
    ap.add_argument("--start", default="2016-01-01")
    ap.add_argument("--end", default="2026-12-31")
    ap.add_argument(
        "--forward-days",
        type=int,
        nargs="+",
        default=[5, 10, 20, 40],
        help="Forward return horizons in trading days",
    )
    ap.add_argument("--ma-window", type=int, default=60, help="SMA window for VXX and VX1")
    ap.add_argument("--bins", type=int, default=5, help="Quantile bins per signal")
    ap.add_argument(
        "--grid-roll",
        type=float,
        nargs="*",
        default=[0.0, 0.05, 0.07, 0.09, 0.11],
        help="Roll thresholds for 2D grid (VX2/VX1−1)",
    )
    ap.add_argument(
        "--grid-elev",
        type=float,
        nargs="*",
        default=[0.0, 0.02, 0.05, 0.08],
        help="VX1 vs MA thresholds for 2D grid",
    )
    ap.add_argument(
        "--use-vxx-ma",
        action="store_true",
        help="2D grid uses vxx_vs_ma instead of vx1_elev_vs_ma on Y axis",
    )
    ap.add_argument("--out-panel", type=Path, default=None, help="Save daily panel CSV")
    ap.add_argument("--out-csv", type=Path, default=None, help="Save combined summary CSV")
    args = ap.parse_args()

    forward_days = sorted(set(args.forward_days))
    print(f"Building panel  {args.start} → {args.end}  MA={args.ma_window}  fwd={forward_days}", flush=True)
    df = build_panel(args.start, args.end, ma_window=args.ma_window, forward_days=forward_days)

    if args.out_panel:
        args.out_panel.parent.mkdir(parents=True, exist_ok=True)
        df.to_csv(args.out_panel, date_format="%Y-%m-%d")
        print(f"Daily panel → {args.out_panel}")

    usable = df.dropna(subset=["vxx_close", "roll_cost_m1_m2"])
    print(
        f"Rows: {len(df):,}  usable: {len(usable):,}  "
        f"VXX median vs MA: {usable['vxx_vs_ma'].median()*100:.2f}%  "
        f"roll median: {usable['roll_cost_m1_m2'].median()*100:.2f}%",
        flush=True,
    )

    all_summaries: list[pd.DataFrame] = []
    elev_col = "vxx_vs_ma" if args.use_vxx_ma else "vx1_elev_vs_ma"
    elev_name = "VXX vs MA" if args.use_vxx_ma else "VX1 vs MA"

    for n in forward_days:
        print_section(f"FORWARD {n} TRADING DAYS — quantile bins")
        qdf = summarize_quantile_bins(df, n, n_bins=args.bins)
        if not qdf.empty:
            print(qdf.to_string(index=False))
            qdf["summary_type"] = "quantile"
            all_summaries.append(qdf)

        print_section(f"FORWARD {n}d — 2D grid: roll × {elev_name}")
        gdf = summarize_2d_grid(
            df,
            n,
            elev_col=elev_col,
            roll_edges=list(args.grid_roll),
            elev_edges=list(args.grid_elev),
        )
        if not gdf.empty:
            cols = ["roll_label", "elev_label", "n", "mean_fwd_pct", "hit_neg_pct", "sharpe_fwd"]
            print(gdf[cols].to_string(index=False))
            gdf["summary_type"] = "grid_2d"
            gdf["elev_axis"] = elev_col
            all_summaries.append(gdf)

        print_section(f"FORWARD {n}d — named combos")
        cdf = summarize_combo_filters(df, n, default_combos())
        print(cdf.to_string(index=False))
        cdf["summary_type"] = "combo"
        all_summaries.append(cdf)

    if all_summaries:
        combined = pd.concat(all_summaries, ignore_index=True)
        out = args.out_csv or LOGS / f"vxx_fwd_return_analysis_{args.start}_{args.end}.csv"
        out = Path(out)
        out.parent.mkdir(parents=True, exist_ok=True)
        combined.to_csv(out, index=False)
        print_section("SAVED")
        print(out)

    print_section("HOW TO READ")
    print(
        "  mean_fwd_pct / hit_neg_pct on VXX spot (yfinance).\n"
        "  Negative mean_fwd_pct + high hit_neg_pct → VXX tended to FALL (good for long puts).\n"
        "  roll = VX2/VX1−1 (monthly roll).  vx1_elev_vs_ma = VX1/SMA−1.\n"
        "  Re-run with --forward-days 5 10 20 and --grid-roll / --grid-elev to explore."
    )


if __name__ == "__main__":
    main()
