#!/usr/bin/env python3
"""
Sweep **stop rules** on ``top10_dual_product_monthly`` (S&P 500 MA slope top-N).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_sp500_stop_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.equity_universe_loaders import load_equity_panel_dict
from RenTech.strategy_stack.ma_slope_cross_sectional import (
    MaSlopeCrossSectional,
    MaSlopeCrossSectionalConfig,
)
from RenTech.strategy_stack.ma_slope_engine import metrics_from_returns

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


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


def _yearly_max_dd(r: pd.Series) -> dict[int, float]:
    out: dict[int, float] = {}
    for year, g in r.groupby(r.index.year):
        eq = (1 + g).cumprod()
        out[int(year)] = float((eq / eq.cummax() - 1).min() * 100)
    return out


def _grid() -> list[MaSlopeCrossSectionalConfig]:
    base = dict(rank_metric="dual_product", rebalance="monthly")
    configs = [MaSlopeCrossSectionalConfig(stop_mode="none", **base)]
    for mult in (2.0, 2.5, 3.0, 3.5):
        configs.append(
            MaSlopeCrossSectionalConfig(stop_mode="atr_trail", atr_multiplier=mult, **base)
        )
    for pct in (0.08, 0.10, 0.12, 0.15, 0.20):
        configs.append(
            MaSlopeCrossSectionalConfig(stop_mode="pct_trail", pct_trail_stop=pct, **base)
        )
    for dd in (0.08, 0.10, 0.12, 0.15):
        configs.append(
            MaSlopeCrossSectionalConfig(stop_mode="portfolio_dd", portfolio_dd_stop=dd, **base)
        )
    return configs


def main() -> None:
    ap = argparse.ArgumentParser()
    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=10)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_sp500_stop_sweep")
    args = ap.parse_args()

    equity_dict = load_equity_panel_dict("sp500", args.yahoo_period)
    grid = _grid()
    rows: list[dict] = []

    for i, cfg in enumerate(grid):
        eng = MaSlopeCrossSectional(config=cfg)
        r_full = eng.generate_returns(equity_dict, top_n=int(args.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))
        y22 = _period_return(r, "2022-01-01", "2022-12-31")
        ydd = _yearly_max_dd(r)
        slug = f"stop_{cfg.stop_mode}"
        if cfg.stop_mode == "atr_trail":
            slug += f"_{cfg.atr_multiplier:g}"
        elif cfg.stop_mode == "pct_trail":
            slug += f"_{int(cfg.pct_trail_stop*100)}pct"
        elif cfg.stop_mode == "portfolio_dd":
            slug += f"_{int(cfg.portfolio_dd_stop*100)}pct"
        rows.append({
            "config_slug": slug,
            **{f"cfg_{k}": v for k, v in asdict(cfg).items()},
            **m,
            "ret_2022_pct": y22,
            "worst_year_dd_pct": min(ydd.values()) if ydd else float("nan"),
            "avg_year_dd_pct": float(np.mean(list(ydd.values()))) if ydd else float("nan"),
        })
        print(f"  [{i+1}/{len(grid)}] {slug}", flush=True)

    df = pd.DataFrame(rows).sort_values("max_dd_pct", ascending=False)
    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    csv_path = Path(f"{prefix}_top{args.top_n}.csv")
    df.to_csv(csv_path, index=False)

    base = df[df.cfg_stop_mode == "none"].iloc[0]
    print(f"\nBaseline: CAGR {base['cagr_pct']:.1f}%  max DD {base['max_dd_pct']:.1f}%  Sharpe {base['sharpe']:.2f}")
    best_dd = df.iloc[0]
    print(f"Best DD:  {best_dd['config_slug']}  CAGR {best_dd['cagr_pct']:.1f}%  max DD {best_dd['max_dd_pct']:.1f}%  Sharpe {best_dd['sharpe']:.2f}")
    sub = df[(df.sharpe >= base.sharpe * 0.85) & (df.max_dd_pct > base.max_dd_pct)]
    if len(sub):
        pick = sub.iloc[0]
        print(f"Balanced: {pick['config_slug']}  CAGR {pick['cagr_pct']:.1f}%  max DD {pick['max_dd_pct']:.1f}%  Sharpe {pick['sharpe']:.2f}")
    print(f"\nWrote {csv_path}")

    meta_path = Path(f"{prefix}_top{args.top_n}_meta.json")
    meta_path.write_text(json.dumps({"csv": str(csv_path), "n_configs": len(grid)}, indent=2) + "\n")


if __name__ == "__main__":
    main()
