#!/usr/bin/env python3
"""
Ride-the-rockets **champ** grid: top 10/12/15 + near-52w-high + 6m-fade kill,
plus ablations (near-only / kill-only).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
      RenTech/strategy_stack/run_ride_rockets_champ_grid.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, fields
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,
    ride_rockets_champ_catalog,
)

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


def _cfg_from_patch(patch: dict) -> Sp500MomentumConfig:
    allowed = {f.name for f in fields(Sp500MomentumConfig)}
    base = asdict(Sp500MomentumConfig())
    cleaned = {k: v for k, v in patch.items() if k in allowed}
    return Sp500MomentumConfig(**{**base, **cleaned})


def _book_stats(rebal: pd.DataFrame) -> dict[str, float]:
    if rebal is None or rebal.empty:
        return {"avg_held": float("nan"), "avg_gross": float("nan"), "n_rebalances": 0}
    live = rebal[rebal["weight"] > 1e-8] if "weight" in rebal.columns else rebal
    g = live.groupby("signal_date")
    held = g.size()
    gross = g["weight"].sum() if "weight" in live.columns else pd.Series(dtype=float)
    n_dates = int(rebal["signal_date"].nunique()) if "signal_date" in rebal.columns else int(len(held))
    return {
        "avg_held": float(held.mean()) if len(held) else float("nan"),
        "avg_gross": float(gross.mean()) if len(gross) else float("nan"),
        "n_rebalances": n_dates,
    }


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("--reuse-equity-cache", action="store_true")
    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=0,
        refresh_cache=not args.reuse_equity_cache,
        pit_universe=True,
        universe="sp500",
    )
    equity_dict = _filter_equity_by_history(equity_dict, min_first)
    print(f"Universe: {len(equity_dict)} names\n", flush=True)

    results: list[dict] = []
    for slug, thesis, patch in ride_rockets_champ_catalog():
        cfg = _cfg_from_patch(patch)
        print(f"=== {slug} ===\n  {thesis}", 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,
        )
        book = _book_stats(rebal)
        row = {
            "variant": slug,
            "thesis": thesis,
            "top_n": cfg.top_n,
            "near_high_min_frac": cfg.near_high_min_frac,
            "require_positive_raw_6": cfg.require_positive_raw_6,
            **{
                k: meta[k]
                for k in (
                    "total_return_pct",
                    "cagr_pct",
                    "sharpe_daily",
                    "max_drawdown_pct",
                    "beta_vs_spy",
                    "corr_vs_spy_daily",
                )
            },
            **book,
        }
        results.append(row)
        print(
            f"  → return {row['total_return_pct']:+.1f}%  CAGR {row['cagr_pct']:+.1f}%  "
            f"Sharpe {row['sharpe_daily']:.2f}  maxDD {row['max_drawdown_pct']:.1f}%  "
            f"avgHeld {row['avg_held']:.1f}  gross {row['avg_gross']:.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")
    df.to_csv(csv_path, index=False)
    Path(f"{prefix}.json").write_text(json.dumps(results, indent=2) + "\n", encoding="utf-8")

    best = df.iloc[0]
    best_ret = df.sort_values("total_return_pct", ascending=False).iloc[0]
    lines = [
        "Ride-the-rockets CHAMP grid",
        f"Window: {args.start} → {args.end}  capital=${args.capital:,.0f}  PIT S&P 500",
        "Recipe: top-N + near-52w-high (≥95%) + 6m-fade kill + monthly",
        f"Command: cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_ride_rockets_champ_grid.py "
        f"--start {args.start} --end {args.end} --reuse-equity-cache",
        "",
        "Also: --ride-rockets-champ on run_sp500_momentum_standard.py (default top 12)",
        "",
        "Ranked by Sharpe:",
        df[
            [
                "variant",
                "total_return_pct",
                "cagr_pct",
                "sharpe_daily",
                "max_drawdown_pct",
                "avg_held",
                "avg_gross",
            ]
        ].to_string(index=False),
        "",
        f"Best Sharpe: {best['variant']}  ({best['thesis']})",
        f"Best return: {best_ret['variant']}  ({best_ret['thesis']})",
        "",
    ]
    metrics_path = Path(f"{prefix}_metrics.txt")
    metrics_path.write_text("\n".join(lines) + "\n", encoding="utf-8")
    print(df.to_string(index=False), flush=True)
    print(f"\nWrote {csv_path}", flush=True)
    print(f"Wrote {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
