#!/usr/bin/env python3
"""
Research runner: SPY overnight close→close filtered by VIX / calmness features.

Builds features at close ``t``, optionally holds SPY until close ``t+1``, and
compares:
  * always-in overnight (every session)
  * hard calm rules (~target 10% coverage)
  * trailing top-frac calm score (default 10%)
  * oracle top-frac realized nights (lookahead ceiling only)

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/run_spy_overnight_vix_calm.py \\
      --start 2016-01-04 --end 2026-06-18 --capital 100000 \\
      --out-prefix RenTech/data/logs/spy_overnight_vix_calm
"""

from __future__ import annotations

import argparse
import json
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.run_qs_top_ideas_backtest import (
    _align,
    _fetch_ohlc,
    _metrics,
)
from RenTech.strategy_stack.spy_overnight_vix_calm import (
    build_feature_frame,
    calm_score,
    conditional_stats,
    feature_lift_table,
    fear_score,
    hard_calm_rules_signal,
    hard_fear_rules_signal,
    next_close_to_close_return,
    oracle_top_frac_mask,
    signal_to_strategy_returns,
    top_frac_score_signal,
)

LOGS = _REPO / "RenTech" / "data" / "logs"
DEFAULT_OUT = LOGS / "spy_overnight_vix_calm"


def _walk_forward_score(
    score: pd.Series,
    fwd: pd.Series,
    *,
    frac: float,
    train_years: int,
    test_years: int,
) -> pd.Series:
    """Expanding calendar walk-forward: fit nothing parametric; re-threshold each fold.

    Uses only trailing score ranks (same as live), but restricts evaluation reporting
    to OOS folds so headline metrics are not full-sample cherry-picked.
    """
    idx = score.index
    sig = pd.Series(False, index=idx)
    if len(idx) < 252:
        return top_frac_score_signal(score, frac=frac)

    start = idx[0]
    end = idx[-1]
    cursor = start + pd.DateOffset(years=train_years)
    while cursor < end:
        test_end = min(cursor + pd.DateOffset(years=test_years), end + pd.Timedelta(days=1))
        fold_idx = idx[(idx >= cursor) & (idx < test_end)]
        # Threshold still uses global trailing window (causal); mark fold for OOS mask
        fold_sig = top_frac_score_signal(score, frac=frac)
        sig.loc[fold_idx] = fold_sig.loc[fold_idx]
        cursor = test_end
    return sig


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("--yahoo-period", default="max")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--top-frac", type=float, default=0.10, help="Target fraction of sessions (~0.10)")
    ap.add_argument("--score-lookback", type=int, default=252)
    ap.add_argument(
        "--mode",
        default="fear_score",
        choices=(
            "calm_score",
            "calm_rules",
            "fear_score",
            "fear_rules",
            "oracle",
        ),
        help="Primary sleeve written to *_daily.csv (default fear_score ≈ top-decile intent)",
    )
    ap.add_argument("--walk-forward", action="store_true", help="Restrict score signal to WF OOS folds")
    ap.add_argument("--wf-train-years", type=int, default=4)
    ap.add_argument("--wf-test-years", type=int, default=1)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end) if str(args.end).strip() else None
    frac = float(args.top_frac)
    cap = float(args.capital)

    # Warmup for SMA200 / percentiles
    fetch_start = start - pd.DateOffset(years=2)

    spy_raw = _fetch_ohlc("SPY", args.yahoo_period)
    vix_raw = _fetch_ohlc("^VIX", args.yahoo_period)
    try:
        vvix_raw = _fetch_ohlc("^VVIX", args.yahoo_period)
    except Exception:
        vvix_raw = None
    try:
        vix3m_raw = _fetch_ohlc("^VIX3M", args.yahoo_period)
    except Exception:
        vix3m_raw = None

    spy_w = _align(spy_raw, fetch_start, end)
    vix_w = _align(vix_raw, fetch_start, end)
    vvix_w = _align(vvix_raw, fetch_start, end) if vvix_raw is not None else None
    vix3m_w = _align(vix3m_raw, fetch_start, end) if vix3m_raw is not None else None

    feats_w = build_feature_frame(spy_w, vix_w, vvix=vvix_w, vix3m=vix3m_w)
    fwd_w = next_close_to_close_return(spy_w)
    calm_w = calm_score(feats_w)
    fear_w = fear_score(feats_w)

    # Restrict evaluation window
    mask_win = spy_w.index >= start
    if end is not None:
        mask_win &= spy_w.index <= end
    spy = spy_w.loc[mask_win]
    feats = feats_w.loc[mask_win]
    fwd = fwd_w.loc[mask_win]
    calm = calm_w.loc[mask_win]
    fear = fear_w.loc[mask_win]

    always = pd.Series(True, index=spy.index)
    calm_rules = hard_calm_rules_signal(feats)
    fear_rules = hard_fear_rules_signal(feats)
    if args.walk_forward:
        calm_sig = _walk_forward_score(
            calm_w,
            fwd_w,
            frac=frac,
            train_years=int(args.wf_train_years),
            test_years=int(args.wf_test_years),
        ).loc[mask_win]
        fear_sig = _walk_forward_score(
            fear_w,
            fwd_w,
            frac=frac,
            train_years=int(args.wf_train_years),
            test_years=int(args.wf_test_years),
        ).loc[mask_win]
    else:
        calm_sig = top_frac_score_signal(
            calm, frac=frac, lookback=int(args.score_lookback)
        )
        fear_sig = top_frac_score_signal(
            fear, frac=frac, lookback=int(args.score_lookback)
        )
    oracle = oracle_top_frac_mask(fwd, frac)

    variants = {
        "always_overnight": always,
        "calm_hard_rules": calm_rules,
        "calm_score_top_frac": calm_sig,
        "fear_hard_rules": fear_rules,
        "fear_score_top_frac": fear_sig,
        "oracle_top_frac_lookahead": oracle,
    }

    ret_map = {k: signal_to_strategy_returns(sig, fwd) for k, sig in variants.items()}
    metrics = {k: _metrics(r, cap) for k, r in ret_map.items()}
    cond = {k: conditional_stats(sig, fwd) for k, sig in variants.items()}

    # Hit rate of top-decile nights captured (precision/recall vs oracle)
    capture = {}
    for k, sig in variants.items():
        if k.startswith("oracle"):
            continue
        inter = (sig & oracle).sum()
        capture[k] = {
            "precision_vs_oracle_pct": round(
                float(inter / sig.sum() * 100.0) if sig.sum() else 0.0, 2
            ),
            "recall_vs_oracle_pct": round(
                float(inter / oracle.sum() * 100.0) if oracle.sum() else 0.0, 2
            ),
        }

    lift = feature_lift_table(feats, fwd, top_frac=frac)

    mode_to_key = {
        "calm_score": "calm_score_top_frac",
        "calm_rules": "calm_hard_rules",
        "fear_score": "fear_score_top_frac",
        "fear_rules": "fear_hard_rules",
        "oracle": "oracle_top_frac_lookahead",
    }
    primary_key = mode_to_key[args.mode]
    primary_ret = ret_map[primary_key]
    primary_sig = variants[primary_key]
    primary_score = fear if "fear" in primary_key else calm

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

    daily = pd.DataFrame(
        {
            "date": spy.index.strftime("%Y-%m-%d"),
            "signal": primary_sig.astype(int).values,
            "calm_score": calm.values,
            "fear_score": fear.values,
            "primary_score": primary_score.values,
            "fwd_c2c": fwd.values,
            "daily_ret": primary_ret.values,
            "daily_pnl_usd": (primary_ret * cap).values,
            "equity_usd": (cap * (1.0 + primary_ret).cumprod()).values,
            "vix": feats["vix"].values,
            "vix_rsi14": feats["vix_rsi14"].values,
            "vix_pct_252": feats["vix_pct_252"].values,
        }
    )
    daily.to_csv(daily_path, index=False)
    lift.to_csv(lift_path, index=False)

    meta = {
        "sleeve": "spy_overnight_vix_calm",
        "return_definition": "close_t_to_close_t1",
        "primary_mode": primary_key,
        "top_frac": frac,
        "score_lookback": int(args.score_lookback),
        "walk_forward": bool(args.walk_forward),
        "start": str(start.date()),
        "end": str(end.date()) if end is not None else str(spy.index[-1].date()),
        "capital": cap,
        "n_sessions": int(len(spy)),
        "has_vvix": bool(vvix_w is not None and len(vvix_w) > 100),
        "has_vix3m": bool(vix3m_w is not None and len(vix3m_w) > 100),
        "metrics": metrics,
        "conditional_fwd_stats": cond,
        "oracle_capture": capture,
        "paths": {
            "daily": str(daily_path),
            "feature_lift": str(lift_path),
        },
        "caveat": (
            "In-sample research on Yahoo adjusted closes; oracle_* is lookahead only. "
            "Empirically, top-decile overnight C2C nights coincide with HIGH VIX/VVIX "
            "and weak SPY — not calm regimes. Use fear_* for 'best nights' intent; "
            "calm_* for low-activity risk-on filters (lower DD, weaker edge)."
        ),
    }
    meta_path.write_text(json.dumps(meta, indent=2))

    lines = [
        "SPY overnight VIX research (calm vs fear)",
        f"window={meta['start']}→{meta['end']}  capital={cap:.0f}  top_frac={frac}",
        f"primary={primary_key}  walk_forward={args.walk_forward}",
        "",
        "NOTE: calm ≠ top-decile overnight winners (those are stress nights).",
        "",
        "--- strategy metrics (exit-day equity) ---",
    ]
    for k, m in metrics.items():
        lines.append(
            f"{k:28s}  ret={m['total_return_pct']:>7.1f}%  "
            f"Sharpe={m['sharpe']:>5.2f}  DD={m['max_dd_pct']:>6.1f}%  "
            f"invested={m.get('pct_days_invested', float('nan')):>5.1f}%"
        )
    lines.append("")
    lines.append("--- conditional forward C2C (signal-day) ---")
    for k, c in cond.items():
        lines.append(
            f"{k:28s}  n={c['n']:>5.0f}  cover={c['coverage_pct']:>5.1f}%  "
            f"mean={c['mean_bps']:>6.1f}bps  hit={c['hit_rate_pct']:>5.1f}%  "
            f"vs_all={c['mean_vs_all_bps']:>+6.1f}bps"
        )
    lines.append("")
    lines.append("--- capture of oracle top-frac nights ---")
    for k, c in capture.items():
        lines.append(
            f"{k:28s}  precision={c['precision_vs_oracle_pct']:>5.1f}%  "
            f"recall={c['recall_vs_oracle_pct']:>5.1f}%"
        )
    lines.append("")
    lines.append(f"wrote {daily_path}")
    lines.append(f"wrote {lift_path}")
    lines.append(f"wrote {meta_path}")
    text = "\n".join(lines) + "\n"
    metrics_path.write_text(text)
    print(text, end="")


if __name__ == "__main__":
    main()
