#!/usr/bin/env python3
"""
Benchmark **MA slope top-N** on Alpaca 1-minute parquet: **5-minute bars vs daily bars**.

Uses the same dual-EMA slope rank (EMA10/EMA50, lookbacks 10/5) and monthly top-10
rebalance as ``run_ma_slope_sp500_topn_standard.py``, but measures slopes on
session-aligned 5-minute bars instead of daily closes.

**Important:** ``fast_period=10`` on 5m bars ≈ 50 minutes of trading, not 10 days.

Example::

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

SPY-only quick test::

    PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_alpaca_5min_benchmark.py \\
        --symbols SPY --top-n 1
"""

from __future__ import annotations

import argparse
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.alpaca_minute_loader import (
    DEFAULT_ALPACA_RTH_DIR,
    bars_per_year,
    compound_intraday_to_daily,
    list_parquet_symbols,
    load_equity_panels,
    load_symbol_bars,
)
from RenTech.strategy_stack.ma_slope_cross_sectional import (
    MaSlopeCrossSectional,
    MaSlopeCrossSectionalConfig,
)

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


def _metrics(
    r: pd.Series,
    *,
    capital: float = 100_000.0,
    ann_factor: float = 252.0,
    label: str = "session",
) -> dict:
    r = r.astype(np.float64).dropna()
    if len(r) < 2:
        return {"label": label}
    eq = capital * (1.0 + r).cumprod()
    n = len(r)
    years = n / ann_factor
    end = float(eq.iloc[-1])
    tot = end / capital - 1.0
    cagr = (end / capital) ** (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(ann_factor)) if sd > 1e-12 else float("nan")
    win = float((r > 0).mean() * 100.0)
    return {
        "label": label,
        "n_bars": int(n),
        "total_return_pct": float(tot * 100.0),
        "cagr_pct": float(cagr * 100.0),
        "max_dd_pct": float(dd * 100.0),
        "sharpe": sharpe,
        "daily_win_rate_pct": win,
        "end_equity_usd": end,
    }


def _yearly(r: pd.Series) -> pd.DataFrame:
    rows = []
    for yr, g in r.groupby(r.index.year):
        eq = (1.0 + g).cumprod()
        rows.append(
            {
                "year": int(yr),
                "return_pct": float((eq.iloc[-1] - 1.0) * 100.0),
                "max_dd_pct": float((eq / eq.cummax() - 1.0).min() * 100.0),
                "n": len(g),
            }
        )
    return pd.DataFrame(rows)


def _run_topn(
    equity_dict: dict[str, pd.DataFrame],
    *,
    top_n: int,
    bars_per_year_val: int,
    stop_mode: str,
    atr_multiplier: float,
) -> pd.Series:
    cfg = MaSlopeCrossSectionalConfig(
        fast_period=10,
        slow_period=50,
        fast_lookback=10,
        slow_lookback=5,
        entry_slope_min=0.0,
        price_above_ma=True,
        rank_metric="dual_product",
        rebalance="monthly",
        stop_mode=stop_mode,  # type: ignore[arg-type]
        atr_multiplier=atr_multiplier,
        bars_per_year=bars_per_year_val,
    )
    eng = MaSlopeCrossSectional(config=cfg)
    return eng.generate_returns(equity_dict, top_n=top_n, verbose=True)


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("--bar-minutes", type=int, default=5)
    ap.add_argument("--top-n", type=int, default=10)
    ap.add_argument("--capital", type=float, default=100_000.0)
    ap.add_argument("--max-tickers", type=int, default=0, help="Cap universe (0 = all parquet symbols)")
    ap.add_argument("--symbols", default="", help="Comma-separated override (e.g. SPY,QQQ,AAPL)")
    ap.add_argument("--no-stops", action="store_true")
    ap.add_argument("--atr-multiplier", type=float, default=2.0)
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    if not data_dir.is_dir():
        raise SystemExit(f"Missing Alpaca data dir: {data_dir}")

    if args.symbols.strip():
        symbols = [s.strip().upper() for s in args.symbols.split(",") if s.strip()]
    else:
        symbols = list_parquet_symbols(data_dir)
        if int(args.max_tickers) > 0:
            symbols = symbols[: int(args.max_tickers)]

    print(f"Loading {len(symbols)} symbols from {data_dir} …", flush=True)
    intra_dict, daily_dict = load_equity_panels(
        symbols,
        data_dir=data_dir,
        bar_minutes=int(args.bar_minutes),
        start=args.start,
        end=args.end,
    )
    if len(intra_dict) < max(1, int(args.top_n)):
        raise SystemExit(f"Only {len(intra_dict)} valid symbols; need at least top_n={args.top_n}")

    print(f"Universe: {len(intra_dict)} names with 5m + daily bars", flush=True)

    bpy = bars_per_year(bar_minutes=int(args.bar_minutes))
    stop_mode = "none" if args.no_stops else "atr_trail"

    print("\n=== Daily-bar MA slope (Alpaca RTH → daily OHLC) ===", flush=True)
    r_daily = _run_topn(
        daily_dict,
        top_n=int(args.top_n),
        bars_per_year_val=252,
        stop_mode=stop_mode,
        atr_multiplier=float(args.atr_multiplier),
    )
    r_daily = r_daily.loc[r_daily.index >= pd.Timestamp(args.start)]
    if args.end.strip():
        r_daily = r_daily.loc[r_daily.index <= pd.Timestamp(args.end)]

    print(f"\n=== {args.bar_minutes}-minute-bar MA slope (same bar-count params) ===", flush=True)
    r_intra = _run_topn(
        intra_dict,
        top_n=int(args.top_n),
        bars_per_year_val=bpy,
        stop_mode=stop_mode,
        atr_multiplier=float(args.atr_multiplier),
    )
    r_intra = r_intra.loc[r_intra.index >= pd.Timestamp(args.start)]
    if args.end.strip():
        r_intra = r_intra.loc[r_intra.index <= pd.Timestamp(args.end)]

    r_intra_daily = compound_intraday_to_daily(r_intra)
    r_intra_daily = r_intra_daily.loc[r_intra_daily.index >= pd.Timestamp(args.start).normalize()]
    if args.end.strip():
        r_intra_daily = r_intra_daily.loc[r_intra_daily.index <= pd.Timestamp(args.end).normalize()]

    # SPY buy-and-hold on same daily calendar (from parquet)
    spy_daily_ret = pd.Series(dtype=float)
    if "SPY" in daily_dict:
        spy_daily_ret = daily_dict["SPY"]["ret"].copy()
        spy_daily_ret.index = pd.to_datetime(spy_daily_ret.index).tz_localize(None)
        spy_daily_ret = spy_daily_ret.loc[r_daily.index.intersection(spy_daily_ret.index)]

    m_daily = _metrics(r_daily, capital=args.capital, ann_factor=252.0, label="daily_bars")
    m_intra_native = _metrics(r_intra, capital=args.capital, ann_factor=float(bpy), label=f"{args.bar_minutes}m_native")
    m_intra_sess = _metrics(r_intra_daily, capital=args.capital, ann_factor=252.0, label=f"{args.bar_minutes}m_compounded_daily")
    m_spy = _metrics(spy_daily_ret, capital=args.capital, ann_factor=252.0, label="SPY_daily") if len(spy_daily_ret) else {}

    print("\n--- Headline comparison (same $100k, same tickers) ---")
    for m in (m_daily, m_intra_sess, m_intra_native, m_spy):
        if not m or "total_return_pct" not in m:
            continue
        print(
            f"  {m['label']:24s}  ret {m['total_return_pct']:+7.1f}%  "
            f"CAGR {m['cagr_pct']:+6.1f}%  Sharpe {m['sharpe']:5.2f}  "
            f"MaxDD {m['max_dd_pct']:6.1f}%  n={m['n_bars']}"
        )

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    slug = f"top{int(args.top_n)}_{int(args.bar_minutes)}m_n{len(intra_dict)}"
    if args.no_stops:
        slug += "_nostops"

    daily_path = Path(f"{prefix}_{slug}_daily_returns.csv")
    intra_path = Path(f"{prefix}_{slug}_intraday_returns.csv")
    sess_path = Path(f"{prefix}_{slug}_intraday_as_daily.csv")
    meta_path = Path(f"{prefix}_{slug}_meta.json")
    metrics_path = Path(f"{prefix}_{slug}_metrics.txt")

    pd.DataFrame({"datetime": r_daily.index, "daily_ret": r_daily.values}).to_csv(daily_path, index=False)
    pd.DataFrame({"datetime": r_intra.index, "bar_ret": r_intra.values}).to_csv(intra_path, index=False)
    pd.DataFrame({"date": r_intra_daily.index, "daily_ret": r_intra_daily.values}).to_csv(sess_path, index=False)

    meta = {
        "window": {"start": args.start, "end": args.end},
        "data_dir": str(data_dir),
        "n_symbols": len(intra_dict),
        "bar_minutes": int(args.bar_minutes),
        "bars_per_year": bpy,
        "top_n": int(args.top_n),
        "config": asdict(
            MaSlopeCrossSectionalConfig(
                stop_mode=stop_mode,  # type: ignore[arg-type]
                bars_per_year=bpy,
            )
        ),
        "metrics": {
            "daily_bars": m_daily,
            "intraday_compounded_daily": m_intra_sess,
            "intraday_native": m_intra_native,
            "spy_daily": m_spy,
        },
        "yearly_daily": _yearly(r_daily).to_dict(orient="records"),
        "yearly_5m_as_daily": _yearly(r_intra_daily).to_dict(orient="records"),
    }
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")

    lines = [
        f"MA slope Alpaca benchmark — {args.bar_minutes}m vs daily",
        f"Window: {args.start} → {args.end}",
        f"Symbols: {len(intra_dict)}  top_n={args.top_n}  stops={stop_mode}",
        f"Params: EMA10/50, lookbacks 10/5, dual_product, monthly rebalance",
        f"Note: 10/50 periods are BAR counts — on 5m ≈ 50min / 4.2hr, not 10/50 days.",
        "",
    ]
    for m in (m_daily, m_intra_sess, m_intra_native, m_spy):
        if "total_return_pct" in m:
            lines.append(
                f"{m['label']}: ret {m['total_return_pct']:+.1f}% CAGR {m['cagr_pct']:+.1f}% "
                f"Sharpe {m['sharpe']:.2f} MaxDD {m['max_dd_pct']:.1f}%"
            )
    metrics_path.write_text("\n".join(lines) + "\n")

    print(f"\nWrote {daily_path}")
    print(f"Wrote {intra_path}")
    print(f"Wrote {sess_path}")
    print(f"Wrote {meta_path}")
    print(f"Wrote {metrics_path}")


if __name__ == "__main__":
    main()
