#!/usr/bin/env python3
"""
Sweep intraday stop rules for **enter 30–90 min, hold MOC** MA slope day-trade.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_alpaca_intraday_stop_sweep.py \\
        --start 2020-01-02 --end 2024-12-31 --top-n 10 --max-tickers 500
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict, replace
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.alpaca_minute_loader import (
    DEFAULT_ALPACA_RTH_DIR,
    compound_intraday_to_daily,
    list_parquet_symbols,
    load_equity_panels,
)
from RenTech.strategy_stack.ma_slope_intraday_daytrade import (
    MaSlopeIntradayDayTrade,
    MaSlopeIntradayDayTradeConfig,
)

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


def _metrics(r_daily: pd.Series) -> dict:
    r = r_daily.astype(np.float64).dropna()
    if len(r) < 2:
        return {}
    eq = (1.0 + r).cumprod()
    years = len(r) / 252.0
    tot = float(eq.iloc[-1] - 1.0)
    cagr = float(eq.iloc[-1] ** (1.0 / years) - 1.0) if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    return {
        "total_return_pct": tot * 100.0,
        "cagr_pct": cagr * 100.0,
        "max_dd_pct": dd * 100.0,
        "sharpe": sharpe,
        "daily_win_rate_pct": float((r > 0).mean() * 100.0),
    }


def _run_cfg(
    intra: dict,
    score_df: pd.DataFrame,
    top_n: int,
    ret_start: pd.Timestamp,
    cfg: MaSlopeIntradayDayTradeConfig,
) -> dict:
    r = MaSlopeIntradayDayTrade(config=cfg).generate_returns(
        intra, top_n=top_n, return_start=ret_start, score_df=score_df, verbose=False
    )
    r_sess = compound_intraday_to_daily(r)
    m = _metrics(r_sess)
    slug = cfg.stop_mode
    if cfg.stop_mode == "atr_trail":
        slug = f"atr{cfg.atr_multiplier:g}x_p{cfg.atr_period}"
    elif cfg.stop_mode == "pct_trail":
        slug = f"pct{int(cfg.pct_trail_stop * 10000) / 100:g}"
    elif cfg.stop_mode == "fixed_pct_entry":
        slug = f"fix{int(cfg.fixed_stop_pct * 10000) / 100:g}pct"
    return {"slug": slug, "config": asdict(cfg), **m}


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--data-dir", type=Path, default=DEFAULT_ALPACA_RTH_DIR)
    ap.add_argument("--start", default="2020-01-02")
    ap.add_argument("--end", default="2024-12-31")
    ap.add_argument("--top-n", type=int, default=10)
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--warmup-sessions", type=int, default=15)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_alpaca_intraday_stop_sweep")
    args = ap.parse_args()

    symbols = list_parquet_symbols(args.data_dir)[: int(args.max_tickers)]
    print(f"Loading {len(symbols)} symbols …", flush=True)
    intra, _ = load_equity_panels(
        symbols,
        data_dir=args.data_dir,
        start=args.start,
        end=args.end,
        warmup_sessions=int(args.warmup_sessions),
    )
    base = MaSlopeIntradayDayTradeConfig()
    ret_start = pd.Timestamp(args.start)
    print("Building slope score panel (once) …", flush=True)
    score_df = MaSlopeIntradayDayTrade(config=base).build_score_panel(intra)

    grid: list[MaSlopeIntradayDayTradeConfig] = [replace(base, stop_mode="none")]
    for mult in (1.5, 2.0, 2.5, 3.0, 3.5):
        grid.append(replace(base, stop_mode="atr_trail", atr_multiplier=mult))
    for pct in (0.01, 0.015, 0.02, 0.025, 0.03, 0.04, 0.05):
        grid.append(replace(base, stop_mode="pct_trail", pct_trail_stop=pct))
    for fix in (0.008, 0.01, 0.012, 0.015, 0.02, 0.025, 0.03):
        grid.append(replace(base, stop_mode="fixed_pct_entry", fixed_stop_pct=fix))

    rows = []
    for i, cfg in enumerate(grid):
        print(f"[{i + 1}/{len(grid)}] {cfg.stop_mode} …", flush=True)
        rows.append(_run_cfg(intra, score_df, int(args.top_n), ret_start, cfg))

    df = pd.DataFrame(rows).sort_values("sharpe", ascending=False)
    baseline = df.loc[df["slug"] == "none"].iloc[0]
    print(
        f"\nBaseline (no stop): ret {baseline['total_return_pct']:+.1f}%  "
        f"Sharpe {baseline['sharpe']:.2f}  MaxDD {baseline['max_dd_pct']:.1f}%"
    )
    print("\nTop 8 by Sharpe:")
    for _, r in df.head(8).iterrows():
        print(
            f"  {r['slug']:16s}  ret {r['total_return_pct']:+7.1f}%  "
            f"Sharpe {r['sharpe']:5.2f}  MaxDD {r['max_dd_pct']:6.1f}%"
        )
    dd_improved = df[df["max_dd_pct"] > baseline["max_dd_pct"]].sort_values("max_dd_pct", ascending=False)
    print("\nBest max-DD improvement (vs baseline) with ret > 50%:")
    ok = dd_improved[dd_improved["total_return_pct"] > 50].sort_values("max_dd_pct", ascending=False)
    for _, r in ok.head(8).iterrows():
        print(
            f"  {r['slug']:16s}  ret {r['total_return_pct']:+7.1f}%  "
            f"Sharpe {r['sharpe']:5.2f}  MaxDD {r['max_dd_pct']:6.1f}%  "
            f"DD Δ {r['max_dd_pct'] - baseline['max_dd_pct']:+.1f}pp"
        )

    slug = f"top{int(args.top_n)}_n{len(intra)}"
    out = args.out_prefix.expanduser().resolve()
    csv_path = Path(f"{out}_{slug}.csv")
    df.to_csv(csv_path, index=False)
    meta = {"baseline": baseline.to_dict(), "best_sharpe": df.iloc[0].to_dict()}
    Path(f"{out}_{slug}_meta.json").write_text(json.dumps(meta, indent=2) + "\n")
    print(f"\nWrote {csv_path}")


if __name__ == "__main__":
    main()
