#!/usr/bin/env python3
"""
Stage 5 — grid-search S&P 500 **MA slope top-N** sleeve.

Sweeps ``rank_metric`` × ``top_n`` × ``rebalance`` (monthly / weekly) on the
full S&P 500 universe. Loads the equity panel once, then scores each config.

Example::

    cd /Users/robzingale/trading_bot
    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_sp500_topn_sweep.py \\
        --start 2016-01-04 --quick

Full grid::

    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_sp500_topn_sweep.py \\
        --start 2016-01-04
"""

from __future__ import annotations

import argparse
import itertools
import json
import sys
from dataclasses import asdict
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.data_loader import DataLoader
from RenTech.strategy_stack.equity_universe_loaders import load_equity_panel_dict
from RenTech.strategy_stack.ma_slope_cross_sectional import (
    MaSlopeCrossSectional,
    MaSlopeCrossSectionalConfig,
)
from RenTech.strategy_stack.main import _compute_daily_backtest_features
from RenTech.strategy_stack.ma_slope_engine import metrics_from_returns

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


def _period_return(r: pd.Series, start: str, end: str) -> float:
    sub = r.loc[start:end]
    if len(sub) < 2:
        return float("nan")
    return float((1.0 + sub).prod() - 1.0) * 100.0


def _stage5_grid(*, quick: bool) -> list[tuple[MaSlopeCrossSectionalConfig, int]]:
    rank_metrics = ["dual_product", "dual_blend"] if quick else [
        "dual_product", "dual_blend", "fast", "slow", "dual_min",
    ]
    top_ns = [5, 10] if quick else [3, 5, 10, 15, 20]
    rebalances = ["monthly"] if quick else ["monthly", "weekly"]
    configs: list[tuple[MaSlopeCrossSectionalConfig, int]] = []
    for metric, top_n, rebal in itertools.product(rank_metrics, top_ns, rebalances):
        cfg = MaSlopeCrossSectionalConfig(
            rank_metric=metric, rebalance=rebal, stop_mode="none",
        )
        configs.append((cfg, int(top_n)))
    return configs


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("--max-tickers", type=int, default=0)
    ap.add_argument("--refresh-cache", action="store_true")
    ap.add_argument("--quick", action="store_true")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    ap.add_argument("--top-n", type=int, default=25, help="Rows in markdown summary")
    args = ap.parse_args()

    equity_dict = load_equity_panel_dict(
        "sp500",
        args.yahoo_period,
        max_tickers=int(args.max_tickers),
        refresh_cache=bool(args.refresh_cache),
    )
    spy_df = _compute_daily_backtest_features(
        DataLoader().fetch_daily("SPY", period=args.yahoo_period)
    )
    spy_r = spy_df["ret"].astype(float)
    spy_r.index = pd.to_datetime(spy_r.index).tz_localize(None)

    grid = _stage5_grid(quick=bool(args.quick))
    rows: list[dict] = []

    for i, (cfg, top_n) in enumerate(grid):
        eng = MaSlopeCrossSectional(config=cfg)
        r_full = eng.generate_returns(equity_dict, top_n=top_n, verbose=False)
        r_full.index = pd.to_datetime(r_full.index).tz_localize(None)
        mask = r_full.index >= pd.Timestamp(args.start)
        if args.end.strip():
            mask &= r_full.index <= pd.Timestamp(args.end)
        r = r_full.loc[mask].fillna(0.0)
        m = metrics_from_returns(r, capital=float(args.capital))
        slug = f"top{top_n}_{cfg.rank_metric}_{cfg.rebalance}"
        row = {
            "config_slug": slug,
            "top_n": top_n,
            **{f"cfg_{k}": v for k, v in asdict(cfg).items()},
            **m,
            "ret_2022_pct": _period_return(r, "2022-01-01", "2022-12-31"),
            "ret_2020_pct": _period_return(r, "2020-01-01", "2020-12-31"),
        }
        rows.append(row)
        if (i + 1) % 5 == 0 or i + 1 == len(grid):
            print(f"  [{i + 1}/{len(grid)}] configs done", flush=True)

    result = pd.DataFrame(rows).sort_values("sharpe", ascending=False)
    slug_out = "stage5" + ("_quick" if args.quick else "")
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    csv_path = Path(f"{prefix}_{slug_out}.csv")
    result.to_csv(csv_path, index=False)

    top = result.head(int(args.top_n))
    cols = ["config_slug", "sharpe", "cagr_pct", "max_dd_pct", "total_return_pct", "ret_2022_pct"]
    md_lines = [
        f"# MA slope S&P 500 top-N sweep — stage 5" + (" (quick)" if args.quick else ""),
        "",
        f"- Window: `{args.start}` → `{args.end or 'latest'}`",
        f"- Universe: {len(equity_dict)} names",
        f"- Configs: **{len(grid)}**",
        f"- CSV: `{csv_path}`",
        "",
        "## Top by Sharpe",
        "",
        "| " + " | ".join(cols) + " |",
        "| " + " | ".join("---" for _ in cols) + " |",
    ]
    for _, row in top.iterrows():
        md_lines.append(
            "| "
            + " | ".join(
                f"{row[c]:.2f}" if isinstance(row[c], (float, np.floating)) else str(row[c])
                for c in cols
            )
            + " |"
        )
    md_path = Path(f"{prefix}_{slug_out}.md")
    md_path.write_text("\n".join(md_lines) + "\n")

    meta_path = Path(f"{prefix}_{slug_out}_meta.json")
    meta_path.write_text(
        json.dumps(
            {
                "stage": 5,
                "quick": args.quick,
                "n_configs": len(grid),
                "n_universe": len(equity_dict),
                "csv": str(csv_path),
                "markdown": str(md_path),
            },
            indent=2,
        )
        + "\n"
    )

    print(f"\nWrote {csv_path}")
    print(f"Wrote {md_path}")
    if len(result):
        best = result.iloc[0]
        print(
            f"\nBest: Sharpe {best['sharpe']:.2f} · CAGR {best['cagr_pct']:.1f}% · "
            f"DD {best['max_dd_pct']:.1f}% · {best['config_slug']}"
        )


if __name__ == "__main__":
    main()
