#!/usr/bin/env python3
"""
Backtest intraday ATR-high breakout long (5m RTH).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_cm_intraday_atr_breakout.py \\
        --preset default --symbols AAPL,MSFT,NVDA,AMZN,GOOGL \\
        --start 2022-01-03 --end 2025-12-31 \\
        --out-prefix RenTech/data/logs/cm_intraday_atr_breakout_mega5
"""

from __future__ import annotations

import argparse
import json
import sys
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.cm_intraday_atr_breakout import (
    atr_breakout_cm_config,
    atr_breakout_config,
    atr_breakout_early_config,
    portfolio_metrics,
    run_backtest,
    trade_stats,
)
from universe_scanner import get_sp100_tickers

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


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("--preset", choices=("default", "early", "cm"), default="default")
    ap.add_argument("--symbols", default="AAPL,MSFT,NVDA,AMZN,GOOGL")
    ap.add_argument("--universe", choices=("custom", "sp100"), default="custom")
    ap.add_argument("--max-tickers", type=int, default=100)
    ap.add_argument("--start", default="2022-01-03")
    ap.add_argument("--end", default="2025-12-31")
    ap.add_argument("--atr-mult", type=float, default=None)
    ap.add_argument("--exit-mode", choices=("moc", "atr"), default=None)
    ap.add_argument("--max-concurrent", type=int, default=10)
    ap.add_argument("--slippage-bps", type=float, default=3.0)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "cm_intraday_atr_breakout")
    args = ap.parse_args()

    have = set(list_parquet_symbols(args.data_dir))
    if args.universe == "sp100":
        syms = [t for t in get_sp100_tickers() if t.upper() in have][: args.max_tickers]
    else:
        syms = [s.strip().upper() for s in args.symbols.split(",") if s.strip()]
    if "SPY" not in syms and "SPY" in have:
        syms = ["SPY"] + syms

    warmup = (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    intra, daily = load_equity_panels(syms, data_dir=args.data_dir, start=warmup, end=args.end)
    spy = intra.pop("SPY", None)
    daily.pop("SPY", None)
    n_book = len(intra)

    cfg = (
        atr_breakout_cm_config()
        if args.preset == "cm"
        else atr_breakout_early_config()
        if args.preset == "early"
        else atr_breakout_config()
    )
    overrides = {"max_concurrent": min(int(args.max_concurrent), n_book), "slippage_bps": float(args.slippage_bps)}
    if args.atr_mult is not None:
        overrides["atr_mult"] = float(args.atr_mult)
    if args.exit_mode is not None:
        overrides["exit_mode"] = args.exit_mode
    cfg = replace(cfg, **overrides)

    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)
    port, trades = run_backtest(intra, daily, cfg=cfg, spy_intra=spy, return_start=ret_start)
    port = port.loc[:end]
    pm = portfolio_metrics(port)
    ts = trade_stats(trades)

    label = f"atr_breakout_{args.preset}"
    print(
        f"{label}  ret {pm.get('total_return_pct', 0):+7.1f}%  "
        f"Sharpe {pm.get('sharpe', 0):5.2f}  DD {pm.get('max_dd_pct', 0):6.1f}%  "
        f"trades={ts.get('n_trades', 0)}  win={ts.get('win_rate_pct', 0):.1f}%  "
        f"avg={ts.get('avg_pnl_pct', 0):+.3f}%",
        flush=True,
    )
    if not trades.empty:
        print(f"  exits: {trades.exit_reason.value_counts().to_dict()}", flush=True)

    out = args.out_prefix.expanduser().resolve()
    row = {"preset": args.preset, "symbols": list(intra.keys()), **pm, **{f"trade_{k}": v for k, v in ts.items()}}
    pd.DataFrame([row]).to_csv(f"{out}.csv", index=False)
    trades.to_csv(f"{out}_trades.csv", index=False)
    compound_intraday_to_daily(port).rename("daily_ret").to_csv(f"{out}_daily.csv")

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_cm_intraday_atr_breakout.py "
        f"--preset {args.preset} --symbols {args.symbols} "
        f"--start {args.start} --end {args.end} --out-prefix {out}"
    )
    meta = {"command": cmd, "config": cfg.__dict__, "metrics": row}
    Path(f"{out}_meta.json").write_text(json.dumps(meta, indent=2, default=str) + "\n")
    Path(f"{out}_metrics.txt").write_text(
        "\n".join([f"command: {cmd}", "", pd.DataFrame([row]).to_string(index=False), ""]) + "\n"
    )
    print(f"\nWrote {out}.csv", flush=True)


if __name__ == "__main__":
    main()
