#!/usr/bin/env python3
"""
Sweep static blends + simple dynamic allocators for **≥10% every calendar year**.

Uses precomputed sleeve daily CSVs ($100k standalone). Metrics focus on:
  * min calendar-year chained return
  * count of years below floor (default 10%)
  * OOS split (train / test) to limit overfit narrative

Dynamic rules are **coarse** (SPY vs SMA200, VIX level) with **quarterly** weight
changes only — no per-day optimization.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 \\
      .venv/bin/python RenTech/strategy_stack/sweep_calendar_year_floor.py \\
      --start 2011-01-03 --end 2025-12-31 --floor 10
"""

from __future__ import annotations

import argparse
import itertools
import json
import math
import sys
from dataclasses import asdict, dataclass
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.combine_best_ideas_stack import _quarter_start_dates  # noqa: E402
from RenTech.strategy_stack.vrp_backtester import load_spy_vix_from_yfinance  # noqa: E402

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

# Canonical sleeve daily artifacts ($100k standalone backtests).
SLEEVE_PATHS: dict[str, Path] = {
    "tactical_aw": LOGS / "tactical_aw_standard_daily.csv",
    "equity_dip": LOGS / "cracking_markets_sp100_dip_daily.csv",
    "tsmom": LOGS / "tsmom_managed_futures_daily.csv",
    "johansen_etf": LOGS / "johansen_triplet_etf_standard_daily.csv",
    "ma_slope_topn": LOGS / "ma_slope_sp500_top10_top10_dual_product_monthly_atr2x_daily.csv",
    "ma_slope_inverse": LOGS
    / "ma_slope_inverse_spy_balanced_bear_dual_or_sma200_spy_slope_or_sma200_invreq_top_n1_SH_daily.csv",
    "vol_edge": LOGS / "volatility_edge_etn_evrp_boc_daily.csv",
    "macro_aw": LOGS / "macro_aw_options_portfolio_eq_daily.csv",
    "lit_vrp": LOGS / "lit_stack_vrp_margin_daily.csv",
    "vxx_regime": LOGS
    / "vxx_regime_mtm_2016_2026_dynamic_vxx_regime_stack_daily_mtm.csv",
}

STOCK_BASE = {
    "tactical_aw": 0.425,
    "vol_edge": 0.11,
    "equity_dip": 0.17,
    "tsmom": 0.0725,
    "johansen_etf": 0.0825,
    "ma_slope_topn": 0.08,
    "ma_slope_inverse": 0.05,
}

OPTIONS_POOL = ("lit_vrp", "vxx_regime", "macro_aw")


@dataclass
class BookMetrics:
    name: str
    min_year_ret: float
    years_below_floor: int
    n_years: int
    cagr_pct: float
    sharpe: float
    max_dd_pct: float
    total_return_pct: float
    worst_years: dict[int, float]
    train_min_year: float
    train_below: int
    test_min_year: float
    test_below: int
    loyo_min: float


def _normalize(w: dict[str, float]) -> dict[str, float]:
    s = sum(w.values())
    if s <= 0:
        raise ValueError("weights sum to zero")
    return {k: v / s for k, v in w.items()}


def _load_sleeve_return(path: Path, idx: pd.DatetimeIndex, *, macro_col: str = "PORTFOLIO_EQUAL_WEIGHT") -> pd.Series:
    df = pd.read_csv(path, parse_dates=["date"])
    df["date"] = pd.to_datetime(df["date"]).dt.normalize()
    df = df.set_index("date").sort_index()
    if "daily_ret" in df.columns:
        ret = df["daily_ret"].astype(np.float64)
    elif "daily_return" in df.columns:
        ret = df["daily_return"].astype(np.float64)
    elif "daily_pnl_usd" in df.columns:
        ret = df["daily_pnl_usd"].astype(np.float64) / 100_000.0
    elif "equity_mtm_usd" in df.columns:
        ret = df["equity_mtm_usd"].astype(np.float64).pct_change().fillna(0.0)
    elif macro_col in df.columns:
        ret = df[macro_col].astype(np.float64).pct_change().fillna(0.0)
    else:
        raise KeyError(f"{path} needs daily_ret, daily_return, daily_pnl_usd, or equity column")
    return ret.reindex(idx).fillna(0.0)


def _load_returns(start: str, end: str) -> tuple[pd.DatetimeIndex, pd.DataFrame]:
    t0, t1 = pd.Timestamp(start), pd.Timestamp(end)
    idx: pd.DatetimeIndex | None = None
    for path in SLEEVE_PATHS.values():
        if not path.is_file():
            continue
        df = pd.read_csv(path, parse_dates=["date"])
        dates = pd.DatetimeIndex(pd.to_datetime(df["date"]).dt.normalize())
        idx = dates if idx is None else idx.union(dates)
    assert idx is not None
    idx = idx[(idx >= t0) & (idx <= t1)].sort_values()
    cols: dict[str, pd.Series] = {}
    for name, path in SLEEVE_PATHS.items():
        if not path.is_file():
            continue
        cols[name] = _load_sleeve_return(path, idx)
    ret = pd.DataFrame(cols, index=idx).fillna(0.0)
    return idx, ret


def _yearly_chained(ret: pd.Series, *, capital: float = 100_000.0) -> pd.Series:
    df = ret.to_frame("r")
    df["year"] = df.index.year
    out: dict[int, float] = {}
    nav = float(capital)
    for yr, g in df.groupby("year", sort=True):
        r_yr = float((1.0 + g["r"]).prod() - 1.0)
        out[int(yr)] = r_yr * 100.0
        nav *= 1.0 + r_yr
    return pd.Series(out)


def _book_metrics(
    port_ret: pd.Series,
    *,
    name: str,
    floor: float,
    train_end: int,
    test_start: int,
) -> BookMetrics:
    yr = _yearly_chained(port_ret)
    full_years = yr[yr.index <= 2025]  # exclude partial 2026
    below = full_years[full_years < floor]
    worst = {int(k): round(float(v), 2) for k, v in full_years.nsmallest(5).items()}

    eq = 100_000.0 * (1.0 + port_ret).cumprod()
    n = len(port_ret)
    years = n / 252.0
    end = float(eq.iloc[-1])
    cagr = (end / 100_000.0) ** (1.0 / years) - 1.0 if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(port_ret.std(ddof=1))
    sharpe = float(port_ret.mean() / sd * math.sqrt(252.0)) if sd > 1e-12 else 0.0

    train_yr = full_years[full_years.index <= train_end]
    test_yr = full_years[full_years.index >= test_start]

    loyo_mins = []
    for y in full_years.index:
        others = full_years[full_years.index != y]
        if len(others):
            loyo_mins.append(float(others.min()))
    loyo_min = min(loyo_mins) if loyo_mins else float(full_years.min())

    return BookMetrics(
        name=name,
        min_year_ret=round(float(full_years.min()), 2),
        years_below_floor=int(len(below)),
        n_years=int(len(full_years)),
        cagr_pct=round(cagr * 100.0, 2),
        sharpe=round(sharpe, 2),
        max_dd_pct=round(dd * 100.0, 2),
        total_return_pct=round((end / 100_000.0 - 1.0) * 100.0, 1),
        worst_years=worst,
        train_min_year=round(float(train_yr.min()), 2) if len(train_yr) else float("nan"),
        train_below=int((train_yr < floor).sum()) if len(train_yr) else 0,
        test_min_year=round(float(test_yr.min()), 2) if len(test_yr) else float("nan"),
        test_below=int((test_yr < floor).sum()) if len(test_yr) else 0,
        loyo_min=round(loyo_min, 2),
    )


def _blend_returns(ret: pd.DataFrame, weights: dict[str, float], fund_scale: float) -> pd.Series:
    """Fund-mode quarterly-sized return series (matches combine ``nav_q`` fund path)."""
    w = _normalize(weights)
    idx = ret.index
    qstarts = _quarter_start_dates(idx)
    nav = 100_000.0
    rets: list[float] = []
    for i, dt in enumerate(idx):
        if i > 0 and dt in qstarts:
            pass  # nav_ref implicit: return is applied to current NAV
        pr = 0.0
        for k, wk in w.items():
            if k in ret.columns:
                pr += float(wk) * float(ret.at[dt, k])
        r = pr * float(fund_scale)
        rets.append(r)
        nav *= 1.0 + r
    return pd.Series(rets, index=idx, dtype=np.float64)


def _static_sweep(
    ret: pd.DataFrame,
    *,
    floor: float,
    fund_scale: float,
    train_end: int,
    test_start: int,
) -> list[BookMetrics]:
    out: list[BookMetrics] = []
    vol_slices = [0.0, 0.10, 0.15, 0.20, 0.25, 0.30, 0.35, 0.40]
    opt_splits = [
        {"lit_vrp": 0.50, "vxx_regime": 0.30, "macro_aw": 0.20},
        {"lit_vrp": 0.40, "vxx_regime": 0.40, "macro_aw": 0.20},
        {"lit_vrp": 0.60, "vxx_regime": 0.25, "macro_aw": 0.15},
        {"lit_vrp": 0.33, "vxx_regime": 0.33, "macro_aw": 0.34},
    ]
    for vol_frac, opt_split in itertools.product(vol_slices, opt_splits):
        stock_frac = 1.0 - vol_frac
        w = {k: v * stock_frac for k, v in STOCK_BASE.items()}
        for ok, ov in opt_split.items():
            if ok in ret.columns:
                w[ok] = w.get(ok, 0.0) + vol_frac * ov
        w = {k: v for k, v in w.items() if v > 1e-9 and k in ret.columns}
        name = f"static_vol{int(vol_frac*100):02d}_" + "+".join(
            f"{k}{int(opt_split[k]*100)}" for k in OPTIONS_POOL if k in ret.columns
        )
        port = _blend_returns(ret, w, fund_scale)
        out.append(
            _book_metrics(
                port,
                name=name,
                floor=floor,
                train_end=train_end,
                test_start=test_start,
            )
        )
    return out


def _regime_weights(
    base: dict[str, float],
    *,
    risk_off: bool,
    vol_boost: bool,
) -> dict[str, float]:
    """Two coarse regimes; shift mass without changing sleeve count."""
    w = dict(base)
    if risk_off:
        for k in ("ma_slope_topn", "equity_dip", "tactical_aw"):
            w[k] = w.get(k, 0.0) * 0.65
        for k in ("tsmom", "macro_aw", "vxx_regime", "vol_edge", "ma_slope_inverse"):
            w[k] = w.get(k, 0.0) * 1.35
    if vol_boost:
        for k in ("lit_vrp", "vxx_regime", "macro_aw", "vol_edge"):
            w[k] = w.get(k, 0.0) * 1.25
        for k in ("ma_slope_topn", "equity_dip"):
            w[k] = w.get(k, 0.0) * 0.85
    return _normalize({k: v for k, v in w.items() if v > 0})


def _dynamic_portfolio(
    ret: pd.DataFrame,
    idx: pd.DatetimeIndex,
    spy: pd.Series,
    vix: pd.Series,
    base_weights: dict[str, float],
    *,
    fund_scale: float,
    sma_days: int = 200,
    vix_hi: float = 22.0,
) -> pd.Series:
    sma = spy.rolling(sma_days, min_periods=sma_days // 2).mean()
    qstarts = _quarter_start_dates(idx)
    cur_w = _normalize(base_weights)
    port = pd.Series(index=idx, dtype=np.float64)
    for i, dt in enumerate(idx):
        if i > 0 and dt in qstarts:
            s = float(spy.reindex([dt]).iloc[0]) if dt in spy.index else float("nan")
            sm = float(sma.reindex([dt]).iloc[0]) if dt in sma.index else float("nan")
            vx = float(vix.reindex([dt]).iloc[0]) if dt in vix.index else float("nan")
            risk_off = math.isfinite(s) and math.isfinite(sm) and s < sm
            vol_boost = math.isfinite(vx) and vx >= vix_hi
            cur_w = _regime_weights(base_weights, risk_off=risk_off, vol_boost=vol_boost)
        r = 0.0
        for k, wk in cur_w.items():
            if k in ret.columns:
                r += float(wk) * float(ret.at[dt, k])
        port.iloc[i] = r * fund_scale
    return port.fillna(0.0)


def _dynamic_sweep(
    ret: pd.DataFrame,
    idx: pd.DatetimeIndex,
    spy: pd.Series,
    vix: pd.Series,
    *,
    floor: float,
    fund_scale: float,
    train_end: int,
    test_start: int,
) -> list[BookMetrics]:
    out: list[BookMetrics] = []
    bases: list[tuple[str, dict[str, float]]] = []
    for vol_frac in (0.15, 0.20, 0.25, 0.30):
        for opt_split in (
            {"lit_vrp": 0.50, "vxx_regime": 0.30, "macro_aw": 0.20},
            {"lit_vrp": 0.40, "vxx_regime": 0.40, "macro_aw": 0.20},
        ):
            stock_frac = 1.0 - vol_frac
            w = {k: v * stock_frac for k, v in STOCK_BASE.items()}
            for ok, ov in opt_split.items():
                if ok in ret.columns:
                    w[ok] = w.get(ok, 0.0) + vol_frac * ov
            w = {k: v for k, v in w.items() if v > 1e-9 and k in ret.columns}
            bases.append((f"base_vol{int(vol_frac*100)}", w))

    for base_name, base_w in bases:
        for sma in (150, 200):
            for vix_hi in (20.0, 22.0, 25.0):
                name = f"dyn_{base_name}_sma{sma}_vix{vix_hi:.0f}"
                port = _dynamic_portfolio(
                    ret,
                    idx,
                    spy,
                    vix,
                    base_w,
                    fund_scale=fund_scale,
                    sma_days=sma,
                    vix_hi=vix_hi,
                )
                out.append(
                    _book_metrics(
                        port,
                        name=name,
                        floor=floor,
                        train_end=train_end,
                        test_start=test_start,
                    )
                )
    return out


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2011-01-03")
    ap.add_argument("--end", default="2025-12-31")
    ap.add_argument("--floor", type=float, default=10.0, help="Min calendar-year return target (%%)")
    ap.add_argument("--max-dd", type=float, default=None, metavar="PCT",
                    help="Optional max drawdown cap (%% positive, e.g. 11 → keep DD ≥ −11%%)")
    ap.add_argument("--fund-scale", type=float, default=1.5)
    ap.add_argument("--train-end", type=int, default=2018, help="Last train year for OOS split")
    ap.add_argument("--test-start", type=int, default=2019, help="First test year")
    ap.add_argument("--top", type=int, default=25)
    ap.add_argument("--out-json", type=Path, default=LOGS / "calendar_year_floor_sweep.json")
    ap.add_argument("--out-md", type=Path, default=LOGS / "calendar_year_floor_sweep.md")
    args = ap.parse_args()

    idx, ret = _load_returns(args.start, args.end)
    print(f"Calendar sweep  {args.start} → {args.end}  ({len(idx)} sessions)", flush=True)
    print(f"Sleeves loaded: {', '.join(ret.columns)}", flush=True)

    yf = load_spy_vix_from_yfinance(
        (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d"),
        (pd.Timestamp(args.end) + pd.Timedelta(days=30)).strftime("%Y-%m-%d"),
    )
    yf.index = pd.to_datetime(yf.index).tz_localize(None).normalize()
    spy = yf["close"].astype(float).reindex(idx).ffill()
    vix = yf["vix_close"].astype(float).reindex(idx).ffill()

    static = _static_sweep(
        ret,
        floor=args.floor,
        fund_scale=args.fund_scale,
        train_end=args.train_end,
        test_start=args.test_start,
    )
    dynamic = _dynamic_sweep(
        ret,
        idx,
        spy,
        vix,
        floor=args.floor,
        fund_scale=args.fund_scale,
        train_end=args.train_end,
        test_start=args.test_start,
    )

    # Baselines
    baselines: list[BookMetrics] = []
    stock_port = _blend_returns(ret, STOCK_BASE, args.fund_scale)
    baselines.append(
        _book_metrics(
            stock_port,
            name="baseline_stock_only_1.5x",
            floor=args.floor,
            train_end=args.train_end,
            test_start=args.test_start,
        )
    )
    if all(k in ret.columns for k in OPTIONS_POOL):
        full_w = dict(STOCK_BASE)
        full_w.update({"lit_vrp": 0.25, "vxx_regime": 0.15, "macro_aw": 0.10})
        # renormalize stock down
        s = sum(STOCK_BASE.values())
        scale = 0.50 / s
        for k in STOCK_BASE:
            full_w[k] = STOCK_BASE[k] * scale
        full_port = _blend_returns(ret, full_w, 2.2)
        baselines.append(
            _book_metrics(
                full_port,
                name="baseline_hybrid_50pct_vol_2.2x",
                floor=args.floor,
                train_end=args.train_end,
                test_start=args.test_start,
            )
        )

    all_rows = baselines + static + dynamic

    def sort_key(m: BookMetrics) -> tuple:
        return (
            m.years_below_floor,
            -m.min_year_ret,
            m.test_below,
            -m.test_min_year,
            -m.cagr_pct,
        )

    ranked = sorted(all_rows, key=sort_key)
    if args.max_dd is not None:
        cap = float(args.max_dd)
        ranked = [m for m in ranked if m.max_dd_pct >= -cap]
        print(f"\nFiltered to max DD ≤ {cap}%: {len(ranked)} configs", flush=True)
    top = ranked[: args.top]

    print(f"\n=== Top {args.top} by (years below {args.floor}%, then min year) ===", flush=True)
    hdr = (
        f"{'name':<42} {'minY':>6} {'<fl':>4} {'train':>6} {'test':>6} "
        f"{'CAGR':>6} {'Sharpe':>6} {'MaxDD':>6}"
    )
    print(hdr, flush=True)
    for m in top:
        print(
            f"{m.name[:42]:<42} {m.min_year_ret:>6.1f} {m.years_below_floor:>4} "
            f"{m.train_min_year:>6.1f} {m.test_min_year:>6.1f} "
            f"{m.cagr_pct:>6.1f} {m.sharpe:>6.2f} {m.max_dd_pct:>6.1f}",
            flush=True,
        )

    passing = [m for m in ranked if m.years_below_floor == 0 and m.n_years >= 10]
    print(f"\nConfigs with ALL years ≥ {args.floor}% (2011-2025): {len(passing)}", flush=True)
    for m in passing[:10]:
        print(f"  {m.name}  min={m.min_year_ret}%  test_min={m.test_min_year}%", flush=True)

    payload = {
        "params": {
            "start": args.start,
            "end": args.end,
            "floor_pct": args.floor,
            "fund_scale": args.fund_scale,
            "train_end": args.train_end,
            "test_start": args.test_start,
        },
        "baselines": [asdict(m) for m in baselines],
        "top": [asdict(m) for m in top],
        "passing_all_years": [asdict(m) for m in passing[:20]],
        "n_static": len(static),
        "n_dynamic": len(dynamic),
    }
    args.out_json.parent.mkdir(parents=True, exist_ok=True)
    args.out_json.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")

    md_lines = [
        f"# Calendar-year floor sweep (≥{args.floor}%/year)",
        "",
        f"Window: **{args.start} → {args.end}** · fund-scale static/dynamic={args.fund_scale} "
        f"(hybrid baseline uses 2.2×)",
        f"OOS split: train ≤ **{args.train_end}**, test ≥ **{args.test_start}**",
        "",
        "## Baselines",
        "",
        "| Config | Min year | Years < floor | Train min | Test min | CAGR | Sharpe | MaxDD |",
        "|--------|----------|---------------|-----------|----------|------|--------|-------|",
    ]
    for m in baselines:
        md_lines.append(
            f"| {m.name} | {m.min_year_ret}% | {m.years_below_floor} | "
            f"{m.train_min_year}% | {m.test_min_year}% | {m.cagr_pct}% | "
            f"{m.sharpe} | {m.max_dd_pct}% |"
        )
    md_lines += [
        "",
        f"## Top {args.top} configs",
        "",
        "| Config | Min year | <floor | Train min | Test min | LOYO min | Worst years |",
        "|--------|----------|--------|-----------|----------|----------|-------------|",
    ]
    for m in top:
        wy = ", ".join(f"{y}:{r}%" for y, r in sorted(m.worst_years.items()))
        md_lines.append(
            f"| {m.name} | {m.min_year_ret}% | {m.years_below_floor} | "
            f"{m.train_min_year}% | {m.test_min_year}% | {m.loyo_min}% | {wy} |"
        )
    md_lines += [
        "",
        f"## Pass floor every year ({len(passing)} configs)",
        "",
    ]
    if passing:
        for m in passing[:15]:
            md_lines.append(
                f"- **{m.name}**: min {m.min_year_ret}%, test min {m.test_min_year}%, "
                f"CAGR {m.cagr_pct}%, Sharpe {m.sharpe}"
            )
    else:
        md_lines.append("_None in this sweep grid._")
    md_lines += [
        "",
        "## Anti-overfit notes",
        "",
        "- Static blends: no tuning on year returns; grid is fixed vol slice × 4 option splits.",
        "- Dynamic rules: **quarterly** rebalance only; 2 binary gates (SPY<SMA, VIX≥threshold).",
        "- Report **train_min** / **test_min** and **LOYO min** (min year when each year excluded once).",
        "- Prefer configs where test_min ≈ train_min and LOYO min stays above ~7%.",
        "- Uses **fund-mode quarterly compounding** (not nav_q stacked-$100k-per-sleeve research combine).",
        "- **nav_q** full Best Ideas clears ≥10%/year (min 2022 ≈ +14.6%) but is a research upper bound —",
        "  not a single-account return; fund-mode 2.2× on the same sleeves still misses the floor in 2022–2024.",
        "",
        f"JSON: `{args.out_json}`",
    ]
    args.out_md.write_text("\n".join(md_lines) + "\n", encoding="utf-8")
    print(f"\nWrote {args.out_json}", flush=True)
    print(f"Wrote {args.out_md}", flush=True)


if __name__ == "__main__":
    main()
