#!/usr/bin/env python3
"""
Sweep S&P 500 momentum **selection filter** variants (baseline + filters 1–7 + combined).

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_sp500_momentum_selection_sweep.py \\
        --start 2016-01-04 --end 2026-04-02 --reuse-equity-cache
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict
from pathlib import Path

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_sp500_dip_standard import _filter_equity_by_history
from RenTech.strategy_stack.sp500_momentum_backtest import load_equity_panel, run_backtest
from RenTech.strategy_stack.sp500_momentum_index import Sp500MomentumConfig, enhanced_selection_defaults

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


def _variants() -> list[tuple[str, dict]]:
    base = asdict(Sp500MomentumConfig())
    rows: list[tuple[str, dict]] = [("A_baseline", {})]
    rows.append(("B_dual_momentum", {
        "require_positive_raw_mom": True,
        "selection_filters": ["dual_momentum"],
    }))
    rows.append(("C_above_sma200", {
        "above_sma_window": 200,
        "selection_filters": ["above_sma"],
    }))
    rows.append(("D_rank_gap5pct", {
        "min_rank_score_pct_gap": 0.05,
        "selection_filters": ["rank_gap"],
    }))
    rows.append(("E_liquidity", {
        "min_market_cap_usd": 1_000_000_000.0,
        "min_adv_usd": 5_000_000.0,
        "selection_filters": ["liquidity"],
    }))
    rows.append(("F_sector_neutral", {
        "sector_neutral_select": True,
        "max_sector_weight": 0.25,
        "selection_filters": ["sector_neutral"],
    }))
    rows.append(("G_residual_mom", {
        "use_residual_momentum": True,
        "selection_filters": ["residual_mom"],
    }))
    rows.append(("H_multi_horizon", {
        "mom_weight_6m": 0.3,
        "mom_weight_12m": 0.7,
        "selection_filters": ["multi_horizon"],
    }))
    rows.append(("I_enhanced_all", enhanced_selection_defaults()))
    rows.append(("J_dual_sma", {
        "require_positive_raw_mom": True,
        "above_sma_window": 200,
        "selection_filters": ["dual_momentum", "above_sma"],
    }))
    out: list[tuple[str, Sp500MomentumConfig]] = []
    for name, patch in rows:
        merged = {**base, **patch}
        out.append((name, Sp500MomentumConfig(**merged)))
    return out


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="2026-04-02")
    ap.add_argument("--yahoo-period", default="max")
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--max-tickers", type=int, default=0)
    ap.add_argument("--reuse-equity-cache", action="store_true", help="Skip Yahoo re-download")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    min_first = pd.Timestamp(args.start).normalize() - pd.Timedelta(days=400)
    equity_dict = load_equity_panel(
        yahoo_period=args.yahoo_period,
        start=args.start,
        end=args.end,
        max_tickers=int(args.max_tickers),
        refresh_cache=not args.reuse_equity_cache,
        pit_universe=True,
    )
    equity_dict = _filter_equity_by_history(equity_dict, min_first)
    print(f"Universe: {len(equity_dict)} names\n", flush=True)

    results: list[dict] = []
    for name, cfg in _variants():
        print(f"=== {name} ===", flush=True)
        _r, _rebal, meta = run_backtest(
            equity_dict=equity_dict,
            cfg=cfg,
            start=args.start,
            end=args.end,
            yahoo_period=args.yahoo_period,
            capital=float(args.capital),
            refresh_shares=False,
            pit_universe=True,
            verbose=True,
        )
        row = {
            "variant": name,
            "filters": cfg.selection_filters,
            **{k: meta[k] for k in (
                "total_return_pct", "cagr_pct", "sharpe_daily",
                "max_drawdown_pct", "beta_vs_spy", "corr_vs_spy_daily",
            )},
        }
        results.append(row)
        print(
            f"  → return {row['total_return_pct']:+.1f}%  Sharpe {row['sharpe_daily']:.2f}  "
            f"maxDD {row['max_drawdown_pct']:.1f}%  β {row['beta_vs_spy']:.2f}\n",
            flush=True,
        )

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    df = pd.DataFrame(results).sort_values("sharpe_daily", ascending=False)
    csv_path = Path(f"{prefix}.csv")
    json_path = Path(f"{prefix}.json")
    df.to_csv(csv_path, index=False)
    json_path.write_text(json.dumps(results, indent=2) + "\n", encoding="utf-8")
    print(df.to_string(index=False), flush=True)
    print(f"\nWrote {csv_path}", flush=True)


if __name__ == "__main__":
    main()
