#!/usr/bin/env python3
"""
Flag **VRP underwater** periods (drawdown from peak), **score overlay JSONLs**
during those days vs calm days, and report **market variables** correlated with
underwater flags (for turning on crisis sleeves).

Uses:
  * ``vrp_backtest_theta.py --export-trades-jsonl`` → ``pnl_usd`` on ``exit_date``
  * ``RenTech/data/vix_futures_cboe.parquet`` for VIX / term structure / VVIX
  * Optional ``yfinance`` for SPY daily returns (skipped if import fails)

Example::

    python RenTech/strategy_stack/vrp_underwater_regime_analysis.py \\
      --vrp-trades RenTech/data/logs/vrp_trades.jsonl \\
      --overlay-dir RenTech/data/logs \\
      --dd-threshold-pct 2.0 \\
      --lead-days 1 5
"""
from __future__ import annotations

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

import numpy as np
import pandas as pd

_REPO = Path(__file__).resolve().parents[2]
LOGS = _REPO / "RenTech/data" / "logs"
VIX_PATH = _REPO / "RenTech" / "data" / "vix_futures_cboe.parquet"

if str(_REPO) not in sys.path:
    sys.path.insert(0, str(_REPO))

from RenTech.strategy_stack.iv_mispricing_complement import (
    _load_jsonl,
    _pnl_series_from_trades,
    _require_jsonl,
)


def _load_vrp_pnl(path: Path) -> pd.Series:
    trades = _load_jsonl(_require_jsonl(path, hint="VRP JSONL"))
    return _pnl_series_from_trades(trades, exit_key="exit_date", pnl_key="pnl_usd")


def _load_overlay_pnl(path: Path) -> pd.Series:
    if not path.is_file():
        return pd.Series(dtype=float)
    return _pnl_series_from_trades(
        _load_jsonl(path), exit_key="exit_date", pnl_key="pnl_total"
    )


def _biserial_correlation(binary: pd.Series, continuous: pd.Series) -> float:
    """Point-biserial r: binary 0/1 vs continuous (aligned, dropna)."""
    m = pd.DataFrame({"b": binary.astype(float), "x": continuous}).dropna()
    if len(m) < 30 or m["b"].std() == 0 or m["x"].std() == 0:
        return float("nan")
    return float(m["b"].corr(m["x"]))


def _build_market_features(vix_df: pd.DataFrame) -> pd.DataFrame:
    """Daily features aligned to vix_df index."""
    out = pd.DataFrame(index=vix_df.index)
    v = vix_df["vix_spot"].astype(float)
    out["vix"] = v
    out["vix_5d_chg"] = v.pct_change(5) * 100.0
    out["vix_20d_chg"] = v.pct_change(20) * 100.0
    if "contango_ratio_ffill" in vix_df.columns:
        out["contango_ratio"] = vix_df["contango_ratio_ffill"].astype(float)
    if "vix3m_vix_ratio" in vix_df.columns:
        out["vix3m_vix"] = vix_df["vix3m_vix_ratio"].astype(float)
    if "vvix" in vix_df.columns:
        out["vvix"] = vix_df["vvix"].astype(float)
    if "vx1_settle" in vix_df.columns and "vx2_settle" in vix_df.columns:
        a = vix_df["vx1_settle"].astype(float)
        b = vix_df["vx2_settle"].astype(float)
        out["vx_calendar_spread"] = (b - a) / a.replace(0, np.nan) * 100.0
    return out


def _optional_spy_returns(idx: pd.DatetimeIndex) -> pd.Series:
    try:
        import yfinance as yf
    except ImportError:
        return pd.Series(dtype=float)
    start = idx.min().strftime("%Y-%m-%d")
    end = (idx.max() + pd.Timedelta(days=1)).strftime("%Y-%m-%d")
    df = yf.download("SPY", start=start, end=end, progress=False, auto_adjust=True)
    if df.empty:
        return pd.Series(dtype=float)
    if isinstance(df.columns, pd.MultiIndex):
        df = df.copy()
        df.columns = df.columns.get_level_values(0)
    close = df["Close"].squeeze()
    close.index = pd.to_datetime(close.index).tz_localize(None).normalize()
    ret = close.pct_change() * 100.0
    return ret.reindex(idx)


def main() -> None:
    ap = argparse.ArgumentParser(
        description="VRP underwater flags, overlay scoring, crisis correlates"
    )
    ap.add_argument("--vrp-trades", type=Path, default=LOGS / "vrp_trades.jsonl")
    ap.add_argument("--start-capital", type=float, default=100_000.0)
    ap.add_argument(
        "--dd-threshold-pct",
        type=float,
        default=2.0,
        help="Flag day as underwater if drawdown from peak <= -this %% (e.g. 2 => -2%%).",
    )
    ap.add_argument(
        "--overlay-dir",
        type=Path,
        default=LOGS,
        help="Directory of *.jsonl overlays (excludes vrp_trades by name).",
    )
    ap.add_argument(
        "--overlay-glob",
        type=str,
        default="*.jsonl",
        help="Glob under overlay-dir for scoring (default all jsonl).",
    )
    ap.add_argument(
        "--skip-files",
        nargs="*",
        default=["vrp_trades.jsonl"],
        help="Basenames to skip when scanning overlay dir.",
    )
    ap.add_argument(
        "--vix-parquet",
        type=Path,
        default=VIX_PATH,
        help="CBOE panel (VIX spot, contango, VVIX, …).",
    )
    ap.add_argument(
        "--lead-days",
        type=int,
        nargs="*",
        default=[1, 5],
        help="Lagged features vs *future* underwater (predictive hints).",
    )
    ap.add_argument("--no-spy", action="store_true", help="Do not fetch SPY via yfinance.")
    ap.add_argument(
        "--export-panel",
        type=Path,
        default=None,
        help="Optional CSV: date, equity, dd_pct, underwater, + features (for research).",
    )
    args = ap.parse_args()

    vrp_pnl = _load_vrp_pnl(args.vrp_trades)
    if vrp_pnl.empty:
        print("ERROR: no VRP PnL", file=sys.stderr)
        sys.exit(1)

    vix = pd.read_parquet(args.vix_parquet)
    vix.index = pd.to_datetime(vix.index).normalize()
    vix.index.name = "trade_date"

    # Dense calendar: VRP trade dates union VIX dates in overlap
    d0 = max(vrp_pnl.index.min(), vix.index.min())
    d1 = min(vrp_pnl.index.max(), vix.index.max())
    all_days = pd.date_range(d0, d1, freq="D")

    daily_pnl = vrp_pnl.reindex(all_days, fill_value=0.0)
    equity = float(args.start_capital) + daily_pnl.cumsum()
    peak = equity.cummax()
    dd_pct = (equity - peak) / peak.replace(0, np.nan) * 100.0
    dd_pct = dd_pct.fillna(0.0)

    underwater = (dd_pct <= -abs(args.dd_threshold_pct)).astype(np.int8)
    # NOTE: We avoid calendar spreads in the live book — document for operators.
    # Underwater is defined on **VRP-only** equity (main book), not full portfolio.

    panel = pd.DataFrame(
        {
            "vrp_pnl": daily_pnl,
            "equity": equity,
            "dd_pct": dd_pct,
            "underwater": underwater,
        },
        index=all_days,
    )

    feats = _build_market_features(vix)
    panel = panel.join(feats.reindex(all_days), how="left")

    if not args.no_spy:
        spy_ret = _optional_spy_returns(all_days)
        if not spy_ret.empty:
            panel["spy_ret_1d"] = spy_ret.reindex(all_days)
            panel["spy_5d_ret"] = panel["spy_ret_1d"].rolling(5).sum()
            panel["spy_20d_ret"] = panel["spy_ret_1d"].rolling(20).sum()

    # --- Concurrent correlation: features vs underwater (same day) ---
    print("=" * 88)
    print("VRP UNDERWATER REGIME ANALYSIS")
    print("=" * 88)
    print(f"VRP trades: {args.vrp_trades}")
    print(f"Start capital: ${args.start_capital:,.0f}")
    print(f"Underwater if drawdown <= -{abs(args.dd_threshold_pct):.2f}% from VRP peak")
    print(f"Date range: {d0.date()} → {d1.date()}  ({len(all_days)} calendar days)")
    print()
    u_days = int(underwater.sum())
    print(f"Underwater days: {u_days} ({100.0 * u_days / len(all_days):.1f}%)")
    print(f"Max drawdown (VRP book): {float(dd_pct.min()):.2f}%")
    print()

    # Exclude tautological columns (underwater is defined from dd_pct / equity path).
    _skip_predictors = {"underwater", "vrp_pnl", "equity", "dd_pct"}
    numeric_cols = [c for c in panel.columns if c not in _skip_predictors]
    print("--- Point-biserial correlation: **external** features vs underwater (same day) ---")
    print("    (dd_pct / equity omitted — they define the underwater flag.)")
    print(f"{'Feature':<22} {'r':>8} {'mean|U':>10} {'mean|~U':>10}")
    print("-" * 55)
    for col in numeric_cols:
        if panel[col].isna().all():
            continue
        r = _biserial_correlation(panel["underwater"], panel[col])
        sub = panel.dropna(subset=[col])
        if sub.empty:
            continue
        m_u = float(sub.loc[sub["underwater"] == 1, col].mean())
        m_c = float(sub.loc[sub["underwater"] == 0, col].mean())
        if not math.isfinite(r):
            continue
        print(f"{col:<22} {r:>+8.3f} {m_u:>10.3f} {m_c:>10.3f}")
    print()

    # --- Leading indicators: today's X vs underwater in h days ---
    lead_feats = list(dict.fromkeys(numeric_cols + (["dd_pct"] if "dd_pct" in panel.columns else [])))
    for h in args.lead_days:
        if h <= 0:
            continue
        future_u = panel["underwater"].shift(-h).fillna(0).astype(np.int8)
        print(f"--- Point-biserial: feature today vs underwater in **{h}** days ---")
        print(f"{'Feature':<22} {'r':>8}")
        print("-" * 35)
        for col in lead_feats:
            if panel[col].isna().all():
                continue
            r = _biserial_correlation(future_u, panel[col])
            if math.isfinite(r):
                print(f"{col:<22} {r:>+8.3f}")
        print()

    # --- Overlay scoring ---
    skip = set(args.skip_files)
    paths = sorted(args.overlay_dir.glob(args.overlay_glob))
    paths = [p for p in paths if p.is_file() and p.name not in skip and p.suffix == ".jsonl"]

    print("--- Overlay daily PnL: mean on underwater vs calm days ---")
    print(
        f"{'File':<38} {'mean|U':>10} {'mean|~U':>10} {'ratio':>8} {'n_U':>6} {'n_tot':>6}"
    )
    print("-" * 88)
    for p in paths:
        o = _load_overlay_pnl(p)
        if o.empty:
            continue
        o = o.reindex(all_days, fill_value=0.0)
        m = panel["underwater"].reindex(all_days, fill_value=0)
        u_mask = m == 1
        mu = float(o[u_mask].mean()) if u_mask.any() else float("nan")
        mc = float(o[~u_mask].mean()) if (~u_mask).any() else float("nan")
        ratio = mu / mc if mc not in (0, float("nan")) and math.isfinite(mc) else float("nan")
        print(
            f"{p.name:<38} {mu:>+10.2f} {mc:>+10.2f} {ratio:>8.2f} "
            f"{int(u_mask.sum()):>6} {len(all_days):>6}"
        )
    print("=" * 88)
    print()
    print(
        "OPERATOR NOTE: Calendar spread overlay was dropped from the live portfolio "
        "(negative edge + DD drag). Do not re-enable without new evidence."
    )
    print()
    print(
        "Interpretation: positive r with underwater means the variable tends to be "
        "**higher** on underwater days (e.g. high VIX). For **crisis-on** rules, "
        "combine 2–3 weak predictors (VIX level, dd depth, SPY drawdown) rather "
        "than a single threshold."
    )

    if args.export_panel is not None:
        p = args.export_panel.expanduser()
        p.parent.mkdir(parents=True, exist_ok=True)
        panel_out = panel.reset_index().rename(columns={"index": "date"})
        panel_out.to_csv(p, index=False)
        print(f"\nWrote panel CSV: {p}")


if __name__ == "__main__":
    main()
