#!/usr/bin/env python3
"""
Canonical **S&P 500 Momentum Index** sleeve (SPMO-style).

Enhanced selection filters (1–7)::

  1. Dual momentum       ``--dual-momentum`` / ``--min-raw-mom``
  2. SMA trend           ``--above-sma 200``
  3. Rank gap            ``--min-rank-gap 0.05``
  4. Liquidity floors    ``--min-mkt-cap`` / ``--min-adv``
  5. Sector-neutral pick ``--sector-neutral-select`` + ``--max-sector-weight``
  6. Residual momentum   ``--residual-momentum``
  7. Multi-horizon blend ``--mom-blend 0.3,0.7``

Preset (all seven): ``--enhanced-selection``

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_sp500_momentum_standard.py \\
        --start 2016-01-04 --end 2026-04-02 --enhanced-selection
"""

from __future__ import annotations

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

_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,
    write_run_artifacts,
)
from RenTech.strategy_stack.sp500_momentum_index import (
    Sp500MomentumConfig,
    best_of_best_defaults,
    enhanced_selection_defaults,
    ride_rockets_champ_defaults,
)

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


def _parse_mom_blend(s: str) -> tuple[float, float]:
    parts = [float(x.strip()) for x in str(s).split(",") if x.strip()]
    if len(parts) != 2:
        raise argparse.ArgumentTypeError("mom-blend must be two comma-separated weights, e.g. 0.3,0.7")
    return parts[0], parts[1]


def _config_from_args(args: argparse.Namespace) -> Sp500MomentumConfig:
    if args.legacy_simple:
        return Sp500MomentumConfig(
            top_n=100 if int(args.top_n) == 0 else int(args.top_n),
            cap_proxy="price",
            use_sp_index_rules=False,
            use_buffer_rule=False,
        )

    kw: dict = {
        "top_n": int(args.top_n),
        "rebalance": args.rebalance,
        "cap_proxy": args.cap_proxy,
        "cash_annual_yield": float(args.cash_yield),
        "use_sp_index_rules": bool(args.sp_index_rules),
        "use_buffer_rule": not args.no_buffer,
        "max_sector_weight": float(args.max_sector_weight),
    }

    if args.enhanced_selection:
        kw.update(enhanced_selection_defaults())
    elif args.ride_rockets_champ:
        # CLI default for --concentrated-top-n is 25; champ wants 15 unless user overrides.
        top = int(args.concentrated_top_n)
        if top == 25 and "--concentrated-top-n" not in sys.argv:
            top = 15
        kw.update(ride_rockets_champ_defaults(top_n=top))
        if float(args.score_weight_power) != 1.0:
            kw["score_weight_power"] = float(args.score_weight_power)
    elif args.best_of_best:
        top = int(args.concentrated_top_n) if int(args.concentrated_top_n) > 0 else 25
        kw.update(best_of_best_defaults(top_n=top))
        if float(args.score_weight_power) != 1.0:
            kw["score_weight_power"] = float(args.score_weight_power)
    else:
        filters: list[str] = []
        if args.dual_momentum or args.min_raw_mom > 0:
            kw["require_positive_raw_mom"] = bool(args.dual_momentum)
            kw["min_raw_mom"] = float(args.min_raw_mom)
            filters.append("dual_momentum")
        if int(args.above_sma) > 0:
            kw["above_sma_window"] = int(args.above_sma)
            filters.append("above_sma")
        if float(args.min_rank_gap) > 0:
            kw["min_rank_score_pct_gap"] = float(args.min_rank_gap)
            filters.append("rank_gap")
        if float(args.min_mkt_cap) > 0 or float(args.min_adv) > 0:
            kw["min_market_cap_usd"] = float(args.min_mkt_cap)
            kw["min_adv_usd"] = float(args.min_adv)
            filters.append("liquidity")
        if args.sector_neutral_select:
            kw["sector_neutral_select"] = True
            filters.append("sector_neutral")
        if args.residual_momentum:
            kw["use_residual_momentum"] = True
            filters.append("residual_mom")
        w6, w12 = args.mom_blend
        if w6 > 0 and w12 > 0 and (abs(w6 - 0.0) > 1e-9 or abs(w12 - 1.0) > 1e-9):
            kw["mom_weight_6m"] = w6
            kw["mom_weight_12m"] = w12
            if w6 > 0:
                filters.append("multi_horizon")
        kw["selection_filters"] = filters

    return Sp500MomentumConfig(**kw)


def _slug_from_config(cfg: Sp500MomentumConfig, *, universe: str, pit: bool, legacy: bool) -> str:
    if legacy:
        base = f"top{cfg.top_n or 100}_legacy_{cfg.rebalance}_{cfg.cap_proxy}"
        return base
    if cfg.selection_filters:
        if "champ" in cfg.selection_filters:
            tag = f"champ{cfg.top_n}"
        elif "best_of_best" not in cfg.selection_filters and len(cfg.selection_filters) >= 5:
            tag = "enhanced"
        elif "best_of_best" in cfg.selection_filters:
            tag = f"bob{cfg.top_n}"
        else:
            tag = "_".join(cfg.selection_filters[:3])
    else:
        tag = "base"
    u = universe.lower().strip()
    uni_tag = "pit" if pit else ("r3k" if u == "russell3000" else "snap")
    sel = f"top{cfg.top_n}" if cfg.top_n > 0 else "quintile"
    return f"{sel}_{tag}_{uni_tag}_{cfg.rebalance}_{cfg.cap_proxy}"


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("--top-n", type=int, default=0)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--max-tickers", type=int, default=0)
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--cash-yield", type=float, default=0.04)
    ap.add_argument("--rebalance", default="semiannual", choices=("semiannual", "monthly"))
    ap.add_argument("--cap-proxy", default="mkt_cap", choices=("mkt_cap", "price", "equal"))
    ap.add_argument("--pit-universe", action=argparse.BooleanOptionalAction, default=True)
    ap.add_argument("--sp-index-rules", action=argparse.BooleanOptionalAction, default=True)
    ap.add_argument("--no-buffer", action="store_true")
    ap.add_argument("--legacy-simple", action="store_true")
    ap.add_argument("--refresh-shares", action="store_true")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT_PREFIX)
    ap.add_argument(
        "--universe",
        default="sp500",
        choices=("sp500", "russell3000"),
        help="Equity universe (default sp500; russell3000 = ~2.5k snapshot)",
    )

    # Enhanced selection 1–7
    ap.add_argument("--enhanced-selection", action="store_true", help="Enable all seven selection filters")
    ap.add_argument(
        "--ride-rockets-champ",
        action="store_true",
        help="Champ recipe: top-N + near-52w-high + 6m-fade kill + monthly (default top 15)",
    )
    ap.add_argument(
        "--best-of-best",
        action="store_true",
        help="Concentrated top-N: dual mom + SMA200 + residual; default top 25",
    )
    ap.add_argument(
        "--concentrated-top-n",
        type=int,
        default=25,
        help="Names for --best-of-best / --ride-rockets-champ (champ default 15 if flag set)",
    )
    ap.add_argument("--score-weight-power", type=float, default=1.0, help="Weight ∝ score^p (default 1; bob uses 1.25)")
    ap.add_argument("--dual-momentum", action="store_true", help="(1) Require positive 12−1 raw return")
    ap.add_argument("--min-raw-mom", type=float, default=0.0)
    ap.add_argument("--above-sma", type=int, default=0, help="(2) Require close > SMA(N); e.g. 200")
    ap.add_argument("--min-rank-gap", type=float, default=0.0, help="(3) Min pct gap top-N vs N+1")
    ap.add_argument("--min-mkt-cap", type=float, default=0.0, help="(4) Min market cap USD")
    ap.add_argument("--min-adv", type=float, default=0.0, help="(4) Min 20d avg dollar volume USD")
    ap.add_argument("--sector-neutral-select", action="store_true", help="(5) Rank within GICS sectors")
    ap.add_argument("--max-sector-weight", type=float, default=0.0, help="(5) Cap sector weight (e.g. 0.25)")
    ap.add_argument("--residual-momentum", action="store_true", help="(6) Beta-adjust vs SPY")
    ap.add_argument("--mom-blend", type=_parse_mom_blend, default=(0.0, 1.0), help="(7) 6m,12m weights e.g. 0.3,0.7")
    ap.add_argument("--residual-beta-window", type=int, default=60)
    args = ap.parse_args()

    end_str = args.end.strip() or str(__import__("pandas").Timestamp.today().date())
    u = args.universe.lower().strip()
    pit_universe = not args.legacy_simple and bool(args.pit_universe) and u == "sp500"
    cfg = _config_from_args(args)
    if args.residual_beta_window != 60:
        cfg = Sp500MomentumConfig(**{**asdict(cfg), "residual_beta_window": int(args.residual_beta_window)})

    min_first = __import__("pandas").Timestamp(args.start).normalize() - __import__("pandas").Timedelta(days=400)
    equity_dict = load_equity_panel(
        universe=u,
        yahoo_period=args.yahoo_period,
        start=args.start,
        end=end_str,
        max_tickers=int(args.max_tickers),
        refresh_cache=bool(args.refresh_cache),
        pit_universe=pit_universe,
    )
    n_loaded = len(equity_dict)
    equity_dict = _filter_equity_by_history(equity_dict, min_first)
    if len(equity_dict) < 50:
        raise SystemExit(f"Too few names ({len(equity_dict)} of {n_loaded})")

    r, rebal, meta = run_backtest(
        equity_dict=equity_dict,
        cfg=cfg,
        start=args.start,
        end=end_str,
        yahoo_period=args.yahoo_period,
        capital=float(args.capital),
        refresh_shares=bool(args.refresh_shares),
        pit_universe=pit_universe,
    )

    cmd_parts = [
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python",
        "RenTech/strategy_stack/run_sp500_momentum_standard.py",
        f"--start {args.start} --end {end_str} --yahoo-period {args.yahoo_period}",
        f"--universe {u} --rebalance {args.rebalance} --cap-proxy {args.cap_proxy}",
    ]
    if args.legacy_simple:
        cmd_parts.append("--legacy-simple")
    elif args.enhanced_selection:
        cmd_parts.append("--enhanced-selection")
    elif args.ride_rockets_champ:
        top = int(args.concentrated_top_n)
        if top == 25 and "--concentrated-top-n" not in sys.argv:
            top = 15
        cmd_parts.append(f"--ride-rockets-champ --concentrated-top-n {top}")
    elif args.best_of_best:
        cmd_parts.append(f"--best-of-best --concentrated-top-n {int(args.concentrated_top_n)}")
    cmd = " ".join(cmd_parts)

    slug = _slug_from_config(cfg, universe=u, pit=pit_universe, legacy=bool(args.legacy_simple))
    paths = write_run_artifacts(
        prefix=args.out_prefix,
        slug=slug,
        r=r,
        rebal=rebal,
        meta={
            **meta,
            "strategy": "Sp500MomentumIndex",
            "token": "sp500_momentum",
            "pit_universe": pit_universe,
            "universe": u,
            "capital_usd": float(args.capital),
            "selection_filters": cfg.selection_filters,
        },
        capital=float(args.capital),
        command=cmd,
    )

    print(
        f"\nSP500 momentum: return {meta['total_return_pct']:+.1f}%  "
        f"CAGR {meta['cagr_pct']:+.1f}%  Sharpe {meta['sharpe_daily']:.2f}  "
        f"maxDD {meta['max_drawdown_pct']:.1f}%  "
        f"β(SPY) {meta['beta_vs_spy']:.2f}  ρ(SPY) {meta['corr_vs_spy_daily']:.2f}",
        flush=True,
    )
    print(f"Filters: {cfg.selection_filters or ['base']}", flush=True)
    print(f"Wrote {paths['daily']}", flush=True)


if __name__ == "__main__":
    main()
