#!/usr/bin/env python3
"""
Sweep Zarattini ORB opening-range lengths (paper: 5 / 10 / 15 / 30 / 60 min).

Loads the equity panel once, then re-runs screen+sim for each OR length.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_orb_zarattini_or_sweep.py \\
        --start 2020-01-02 --end 2023-12-29 --capital 25000 \\
        --out-prefix RenTech/data/logs/orb_zarattini_or_sweep_2020_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"
DEFAULT_OR_MINUTES = (5, 10, 15, 30, 60)


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="2023-12-29")
    ap.add_argument("--capital", type=float, default=25_000.0)
    ap.add_argument("--top-n", type=int, default=20)
    ap.add_argument("--max-tickers", type=int, default=0, help="0 = all symbols")
    ap.add_argument("--or-minutes", default="5,10,15,30,60")
    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_or_sweep")
    args = ap.parse_args()

    or_list = [int(x.strip()) for x in str(args.or_minutes).split(",") if x.strip()]
    if not or_list:
        or_list = list(DEFAULT_OR_MINUTES)

    t0 = time.time()
    all_syms = sorted(list_parquet_symbols(args.data_dir))
    syms = all_syms[: args.max_tickers] if args.max_tickers > 0 else all_syms
    if "SPY" not in syms and "SPY" in all_syms:
        syms = ["SPY"] + [s for s in syms if s != "SPY"]
    print(f"universe: {len(syms)} symbols | OR lengths: {or_list}", 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,
        bar_minutes=5,
        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)

    base = replace(
        zarattini_config(),
        starting_capital=float(args.capital),
        top_n=int(args.top_n),
        slippage_bps=float(args.slippage_bps),
        commission_per_share=float(args.commission),
        bar_minutes=5,
    )

    rows: list[dict] = []
    args.out_prefix.parent.mkdir(parents=True, exist_ok=True)

    for or_m in or_list:
        t1 = time.time()
        cfg = replace(base, or_minutes=int(or_m))
        print(f"\n=== OR {or_m}m ===", flush=True)
        port_r, trades, meta = run_orb_backtest(
            intra,
            daily,
            cfg=cfg,
            return_start=pd.Timestamp(args.start),
            return_end=pd.Timestamp(args.end),
            verbose=True,
        )
        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["elapsed_s"] = float(time.time() - t1)

        print(
            f"ORB {or_m:>2}m  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"β {meta.get('beta_spy', float('nan')):5.2f}  "
            f"trades={meta.get('n_trades', 0)}  "
            f"end=${meta.get('ending_equity', 0):,.0f}  "
            f"({meta['elapsed_s']:.0f}s)",
            flush=True,
        )

        stem = f"{args.out_prefix}_or{or_m}"
        port_r.to_csv(f"{stem}_daily.csv", header=["daily_return"])
        if not trades.empty:
            trades.to_csv(f"{stem}_trades.csv", index=False)
        Path(f"{stem}_meta.json").write_text(json.dumps(meta, indent=2), encoding="utf-8")

        rows.append(
            {
                "or_minutes": or_m,
                "total_return_pct": meta.get("total_return_pct"),
                "cagr_pct": meta.get("cagr_pct"),
                "sharpe": meta.get("sharpe"),
                "max_dd_pct": meta.get("max_dd_pct"),
                "beta_spy": meta.get("beta_spy"),
                "n_trades": meta.get("n_trades"),
                "win_rate_pct": meta.get("win_rate_pct"),
                "avg_pnl_pct": meta.get("avg_pnl_pct"),
                "ending_equity": meta.get("ending_equity"),
                "n_sessions_traded": meta.get("n_sessions_traded"),
                "elapsed_s": meta.get("elapsed_s"),
            }
        )

    summary = pd.DataFrame(rows).sort_values("or_minutes")
    summary_path = Path(f"{args.out_prefix}_summary.csv")
    summary.to_csv(summary_path, index=False)

    print("\n=== OR length sweep summary ===", flush=True)
    print(summary.to_string(index=False), flush=True)
    best = summary.sort_values("sharpe", ascending=False).iloc[0]
    print(
        f"\nBest by Sharpe: OR {int(best['or_minutes'])}m "
        f"(Sharpe {best['sharpe']:.2f}, ret {best['total_return_pct']:+.1f}%, "
        f"CAGR {best['cagr_pct']:+.1f}%)",
        flush=True,
    )
    print(f"summary -> {summary_path}", flush=True)
    print(f"total elapsed {time.time() - t0:.1f}s", flush=True)

    metrics_path = Path(f"{args.out_prefix}_metrics.txt")
    metrics_path.write_text(
        f"command: {' '.join(sys.argv)}\n"
        f"window: {args.start} → {args.end}\n"
        f"capital: {args.capital}\n"
        f"n_symbols: {len(intra)}\n"
        f"summary:\n{summary.to_string(index=False)}\n"
        f"best_sharpe_or_minutes: {int(best['or_minutes'])}\n",
        encoding="utf-8",
    )
    print(f"metrics -> {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
