#!/usr/bin/env python3
"""
Backtest three intraday CM dip variants (A/B/C) on Alpaca 5m RTH parquet.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_cm_intraday_dip.py \\
        --universe sp100 --start 2022-01-03 --end 2025-12-31 \\
        --out-prefix RenTech/data/logs/cm_intraday_dip
"""

from __future__ import annotations

import argparse
import json
import sys
import time
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_dip import (
    CmIntradayConfig,
    portfolio_metrics,
    run_backtest,
    trade_stats,
)
from universe_scanner import get_sp100_tickers

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


def _resolve_symbols(universe: str, data_dir: Path, max_tickers: int) -> list[str]:
    have = set(list_parquet_symbols(data_dir))
    if universe == "sp100":
        want = [t.upper() for t in get_sp100_tickers(universe="sp100")]
    elif universe == "parquet":
        want = sorted(have)
    else:
        raise ValueError(f"unknown universe {universe}")
    syms = [t for t in want if t in have]
    if "SPY" not in syms and "SPY" in have:
        syms.insert(0, "SPY")
    return syms[: int(max_tickers)]


def _yearly(port_r: pd.Series) -> list[dict]:
    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": int(len(grp))})
    return rows


def _run_one(
    variant: str,
    intra: dict,
    daily: dict,
    spy_intra: pd.DataFrame | None,
    *,
    ret_start: pd.Timestamp,
    end: pd.Timestamp,
    top_n: int,
    max_concurrent: int,
    slippage_bps: float,
) -> dict:
    cfg = CmIntradayConfig(
        variant=variant,
        top_n=int(top_n),
        max_concurrent=int(max_concurrent),
        slippage_bps=float(slippage_bps),
    )
    t0 = time.perf_counter()
    port_r, trades = run_backtest(intra, daily, cfg=cfg, spy_intra=spy_intra, return_start=ret_start)
    port_r = port_r.loc[port_r.index <= end]
    elapsed = time.perf_counter() - t0
    pm = portfolio_metrics(port_r)
    ts = trade_stats(trades)
    return {
        "variant": variant,
        "portfolio": pm,
        "trades": ts,
        "yearly": _yearly(port_r),
        "elapsed_sec": round(elapsed, 1),
        "port_r": port_r,
        "trades_df": trades,
    }


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("--universe", choices=("sp100", "parquet"), default="sp100")
    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("--top-n", type=int, default=10, help="Variant C cross-sectional cap")
    ap.add_argument("--max-concurrent", type=int, default=10, help="Max simultaneous positions (all variants)")
    ap.add_argument("--slippage-bps", type=float, default=3.0)
    ap.add_argument(
        "--variants",
        default="vwap_reclaim,pdl_touch,relative_washout",
        help="Comma-separated subset of A/B/C keys",
    )
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "cm_intraday_dip")
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    symbols = _resolve_symbols(args.universe, data_dir, args.max_tickers)
    if len(symbols) < 20:
        raise SystemExit(f"Too few symbols ({len(symbols)}) in {data_dir}")

    warmup_start = (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    print(f"Loading {len(symbols)} symbols ({args.universe}) {warmup_start} → {args.end} …", flush=True)
    intra, daily = load_equity_panels(
        symbols,
        data_dir=data_dir,
        bar_minutes=5,
        start=warmup_start,
        end=args.end,
        warmup_sessions=0,
    )
    spy_intra = intra.pop("SPY", None)
    daily.pop("SPY", None)
    print(f"  loaded n={len(intra)} (SPY bench={'yes' if spy_intra is not None else 'no'})", flush=True)

    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)
    variants = [v.strip() for v in args.variants.split(",") if v.strip()]
    out_prefix = args.out_prefix.expanduser().resolve()

    summary_rows = []
    meta: dict = {
        "universe": args.universe,
        "n_symbols": len(intra),
        "start": args.start,
        "end": args.end,
        "top_n": args.top_n,
        "max_concurrent": args.max_concurrent,
        "slippage_bps": args.slippage_bps,
        "variants": {},
    }

    for variant in variants:
        print(f"\n=== {variant} ===", flush=True)
        res = _run_one(
            variant,
            intra,
            daily,
            spy_intra,
            ret_start=ret_start,
            end=end,
            top_n=args.top_n,
            max_concurrent=args.max_concurrent,
            slippage_bps=args.slippage_bps,
        )
        pm, ts = res["portfolio"], res["trades"]
        summary_rows.append(
            {
                "variant": variant,
                **pm,
                **{f"trade_{k}": v for k, v in ts.items()},
                "elapsed_sec": res["elapsed_sec"],
            }
        )
        meta["variants"][variant] = {
            "portfolio": pm,
            "trades": ts,
            "yearly": res["yearly"],
            "elapsed_sec": res["elapsed_sec"],
        }
        print(
            f"  ret {pm.get('total_return_pct', float('nan')):+7.1f}%  "
            f"CAGR {pm.get('cagr_pct', float('nan')):+6.1f}%  "
            f"Sharpe {pm.get('sharpe', float('nan')):5.2f}  "
            f"DD {pm.get('max_dd_pct', float('nan')):6.1f}%  "
            f"trades={ts.get('n_trades', 0)}  "
            f"win={ts.get('win_rate_pct', float('nan')):.1f}%  "
            f"avg={ts.get('avg_pnl_pct', float('nan')):+.2f}%",
            flush=True,
        )
        for y in res["yearly"]:
            print(f"    {y['year']}: {y['return_pct']:+.1f}%", flush=True)

        trades_path = Path(f"{out_prefix}_{variant}_trades.csv")
        res["trades_df"].to_csv(trades_path, index=False)
        daily_path = Path(f"{out_prefix}_{variant}_daily.csv")
        compound_intraday_to_daily(res["port_r"]).rename("daily_ret").to_csv(daily_path)

    summary = pd.DataFrame(summary_rows)
    summary_path = Path(f"{out_prefix}_summary.csv")
    summary.to_csv(summary_path, index=False)
    meta_path = Path(f"{out_prefix}_meta.json")
    meta_path.write_text(json.dumps(meta, indent=2, default=str) + "\n")

    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_cm_intraday_dip.py "
        f"--universe {args.universe} --max-tickers {args.max_tickers} "
        f"--start {args.start} --end {args.end} --top-n {args.top_n} "
        f"--max-concurrent {args.max_concurrent} "
        f"--slippage-bps {args.slippage_bps} "
        f"--out-prefix {out_prefix}"
    )
    metrics_path = Path(f"{out_prefix}_metrics.txt")
    metrics_path.write_text(
        "\n".join(
            [
                f"command: {cmd}",
                f"window: {args.start} → {args.end}",
                f"universe: {args.universe} n={len(intra)}",
                f"slippage_bps: {args.slippage_bps}",
                "",
                summary.to_string(index=False),
                "",
            ]
        )
        + "\n"
    )

    print(f"\nWrote {summary_path}", flush=True)
    print(f"Wrote {meta_path}", flush=True)
    print(f"Wrote {metrics_path}", flush=True)


if __name__ == "__main__":
    main()
