#!/usr/bin/env python3
"""
Out-of-sample intraday MA-slope test on **2025–2026** Alpaca RTH data.

Compares baseline vs ``confirm_entry_4b`` (and optional slippage) on the same
~500-name alphabetical universe as the enhancement sweep.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_alpaca_intraday_oos_2025.py \\
        --min-symbols 400 --top-n 10 --max-tickers 500
"""

from __future__ import annotations

import argparse
import json
import sys
import time
from dataclasses import replace
from pathlib import Path

import pandas as pd
import pyarrow.parquet as pq

_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,
    list_parquet_symbols,
    load_equity_panels,
)
from RenTech.strategy_stack.ma_slope_intraday_enhanced import (
    EnhancedIntradayEngine,
    baseline_enhanced_config,
    metrics_daily,
)
from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

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


def _symbols_with_2025(data_dir: Path, symbols: list[str], min_date: str) -> list[str]:
    cutoff = pd.Timestamp(min_date)
    out: list[str] = []
    for s in symbols:
        path = data_dir / f"{s}.parquet"
        if not path.is_file():
            continue
        t = pq.read_table(path, columns=["datetime"])
        mx = pd.Timestamp(t["datetime"].to_pandas().max())
        if mx >= cutoff:
            out.append(s)
    return out


def _yearly(r: pd.Series) -> pd.DataFrame:
    from RenTech.strategy_stack.alpaca_minute_loader import compound_intraday_to_daily

    ds = compound_intraday_to_daily(r).dropna()
    rows = []
    for y, grp in ds.groupby(ds.index.year):
        eq = (1 + grp).cumprod()
        rows.append({"year": int(y), "return_pct": float((eq.iloc[-1] - 1) * 100), "n_days": len(grp)})
    return pd.DataFrame(rows)


def _run_variant(
    name: str,
    cfg,
    intra: dict,
    sector_map: dict,
    *,
    top_n: int,
    ret_start: pd.Timestamp,
    end: pd.Timestamp,
) -> dict:
    eng = EnhancedIntradayEngine(config=cfg)
    panels = eng.build_panels(intra, sector_map)
    port_r = eng.run(intra, top_n, panels=panels, return_start=ret_start)
    port_r = port_r.loc[port_r.index <= end]
    m = metrics_daily(port_r)
    yearly = _yearly(port_r)
    return {"variant": name, "metrics": m, "yearly": yearly.to_dict(orient="records")}


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="2025-01-02")
    ap.add_argument("--end", default="2026-06-18")
    ap.add_argument("--top-n", type=int, default=10)
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--min-symbols", type=int, default=400, help="Wait until this many names have 2025+ data")
    ap.add_argument("--wait-min", type=int, default=0, help="Poll every N minutes for min-symbols (0=run now)")
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_intraday_oos_2025")
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    all_syms = list_parquet_symbols(data_dir)[: int(args.max_tickers)]
    if "SPY" not in all_syms:
        all_syms.insert(0, "SPY")

    sec_df = load_sp500_sectors()
    sector_map = dict(zip(sec_df["ticker"].astype(str).str.upper(), sec_df["sector"].astype(str)))

    ready: list[str] = []
    while True:
        ready = _symbols_with_2025(data_dir, all_syms, args.start)
        print(f"2025+ coverage: {len(ready)}/{len(all_syms)} symbols", flush=True)
        if len(ready) >= int(args.min_symbols) or int(args.wait_min) <= 0:
            break
        print(f"  waiting {args.wait_min}m for download …", flush=True)
        time.sleep(int(args.wait_min) * 60)

    if len(ready) < 50:
        raise SystemExit(f"Too few symbols with 2025 data ({len(ready)}). Wait for Alpaca download.")

    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)

    print(f"Loading {len(ready)} symbols {args.start} → {args.end} …", flush=True)
    intra, _ = load_equity_panels(
        ready, data_dir=data_dir, start="2024-11-01", end=str(end.date()), warmup_sessions=15
    )
    print(f"  panels loaded n={len(intra)}", flush=True)

    base = baseline_enhanced_config()
    variants = [
        ("baseline", base),
        ("confirm_entry_4b", replace(base, hold_mode="confirm_entry", confirm_lag_bars=4)),
        ("confirm_4b_5bps", replace(base, hold_mode="confirm_entry", confirm_lag_bars=4, slippage_bps=5.0)),
    ]

    rows = []
    meta = {"n_symbols_requested": len(ready), "n_symbols_loaded": len(intra), "start": args.start, "end": args.end}
    for name, cfg in variants:
        row = _run_variant(name, cfg, intra, sector_map, top_n=int(args.top_n), ret_start=ret_start, end=end)
        m = row["metrics"]
        rows.append({"variant": name, **m})
        meta[name] = row
        print(
            f"  {name:18s}  ret {m.get('total_return_pct', float('nan')):+7.1f}%  "
            f"Sharpe {m.get('sharpe', float('nan')):5.2f}  DD {m.get('max_dd_pct', float('nan')):6.1f}%",
            flush=True,
        )

    df = pd.DataFrame(rows)
    out = args.out_prefix.expanduser().resolve()
    csv_path = Path(f"{out}_top{args.top_n}.csv")
    df.to_csv(csv_path, index=False)
    Path(f"{out}_top{args.top_n}_meta.json").write_text(json.dumps(meta, indent=2, default=str) + "\n")

    print(f"\nWrote {csv_path}", flush=True)
    for name, _ in variants:
        yr = meta[name]["yearly"]
        print(f"  {name} yearly:", flush=True)
        for y in yr:
            print(f"    {y['year']}: {y['return_pct']:+.1f}%  ({y['n_days']}d)", flush=True)


if __name__ == "__main__":
    main()
