#!/usr/bin/env python3
"""
Backtest Zarattini 5-minute ORB on Stocks in Play (Alpaca 1m RTH parquet).

Paper window: 2016-01-01 → 2023-12-31 · $25k start · top-20 rel-vol names/day.

Example (smoke)::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_orb_zarattini.py \\
        --start 2016-01-04 --end 2016-06-30 --max-tickers 200

Full replication attempt (slow — loads broad universe)::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_orb_zarattini.py \\
        --start 2016-01-04 --end 2023-12-31 --capital 25000 \\
        --out-prefix RenTech/data/logs/orb_zarattini_2016_2023
"""

from __future__ import annotations

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

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.orb_zarattini import beta_vs_spy, run_orb_backtest, zarattini_config

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


def _yearly(port_r: pd.Series) -> pd.DataFrame:
    ds = compound_intraday_to_daily(port_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 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="2016-01-04")
    ap.add_argument("--end", default="2023-12-31")
    ap.add_argument("--capital", type=float, default=25_000.0)
    ap.add_argument("--top-n", type=int, default=20)
    ap.add_argument("--or-minutes", type=int, default=5, help="Opening-range length in minutes")
    ap.add_argument("--max-tickers", type=int, default=0, help="0 = all symbols in data dir")
    ap.add_argument("--slippage-bps", type=float, default=3.0)
    ap.add_argument("--commission", type=float, default=0.0035)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "orb_zarattini")
    args = ap.parse_args()

    t0 = time.time()
    all_syms = sorted(list_parquet_symbols(args.data_dir))
    if args.max_tickers > 0:
        syms = all_syms[: args.max_tickers]
    else:
        syms = all_syms
        if "SPY" not in syms:
            syms = ["SPY"] + syms
    print(f"universe: {len(syms)} symbols from {args.data_dir} | OR={args.or_minutes}m", flush=True)

    warmup = (pd.Timestamp(args.start) - pd.Timedelta(days=60)).strftime("%Y-%m-%d")
    intra, daily = load_equity_panels(
        syms,
        data_dir=args.data_dir,
        start=warmup,
        end=args.end,
        verbose=True,
    )
    spy_daily = pd.Series(dtype=float)
    if "SPY" in daily:
        spy_df = daily.pop("SPY")
        if not spy_df.empty and "ret" in spy_df.columns:
            spy_daily = spy_df["ret"].astype(float)
    if "SPY" in intra:
        spy_ib = intra.pop("SPY")
        if spy_daily.empty and not spy_ib.empty:
            spy_daily = compound_intraday_to_daily(spy_ib["ret"].astype(float))
    print(f"loaded {len(intra)} tradeable symbols", flush=True)

    cfg = replace(
        zarattini_config(),
        starting_capital=float(args.capital),
        top_n=int(args.top_n),
        or_minutes=int(args.or_minutes),
        slippage_bps=float(args.slippage_bps),
        commission_per_share=float(args.commission),
    )

    port_r, trades, meta = run_orb_backtest(
        intra,
        daily,
        cfg=cfg,
        return_start=pd.Timestamp(args.start),
        return_end=pd.Timestamp(args.end),
    )
    meta["beta_spy"] = float(beta_vs_spy(port_r, spy_daily)) if not spy_daily.empty else float("nan")
    meta["window"] = f"{args.start} → {args.end}"
    meta["n_symbols_loaded"] = int(len(intra))
    meta["command"] = " ".join(sys.argv)

    yr = _yearly(port_r)
    elapsed = time.time() - t0

    print(
        f"ORB Zarattini OR={args.or_minutes}m  {meta.get('window')}  "
        f"ret {meta.get('total_return_pct', 0):+8.1f}%  "
        f"CAGR {meta.get('cagr_pct', 0):+6.1f}%  "
        f"Sharpe {meta.get('sharpe', 0):5.2f}  "
        f"DD {meta.get('max_dd_pct', 0):6.1f}%  "
        f"β(SPY) {meta.get('beta_spy', float('nan')):5.2f}  "
        f"trades={meta.get('n_trades', 0)}  "
        f"end=${meta.get('ending_equity', 0):,.0f}",
        flush=True,
    )
    if not trades.empty:
        print(f"  exits: {trades.exit_reason.value_counts().to_dict()}", flush=True)
        print(f"  sides: {trades.variant.value_counts().to_dict()}", flush=True)
    print(f"  screen={meta.get('screen_rows')} picks={meta.get('pick_rows')} sessions={meta.get('n_sessions_traded')}", flush=True)
    print(f"elapsed {elapsed:.1f}s", flush=True)

    args.out_prefix.parent.mkdir(parents=True, exist_ok=True)
    # Best Ideas sleeve unit CSV (date, daily_ret, daily_pnl_usd, equity_usd, margin_usd).
    capital = float(args.capital)
    if not port_r.empty:
        eq = capital * (1.0 + port_r.astype(float)).cumprod()
        sleeve = pd.DataFrame(
            {
                "date": pd.to_datetime(port_r.index).normalize().strftime("%Y-%m-%d"),
                "daily_ret": port_r.astype(float).values,
                "daily_pnl_usd": (port_r.astype(float) * capital).values,
                "equity_usd": eq.values,
                "margin_usd": capital * 0.25,
            }
        )
        sleeve.to_csv(f"{args.out_prefix}_daily.csv", index=False)
    else:
        port_r.to_csv(f"{args.out_prefix}_daily.csv", header=["daily_return"])
    if not trades.empty:
        trades.to_csv(f"{args.out_prefix}_trades.csv", index=False)
    yr.to_csv(f"{args.out_prefix}_yearly.csv", index=False)
    meta_path = Path(f"{args.out_prefix}_meta.json")
    meta_path.write_text(json.dumps(meta, indent=2), encoding="utf-8")
    print(f"meta -> {meta_path}", flush=True)
    print(f"daily -> {args.out_prefix}_daily.csv", flush=True)


if __name__ == "__main__":
    main()
