#!/usr/bin/env python3
"""
Sweep objective vol/ATR filters for ``pdl_touch_short`` on SP100.

Filters use **prior-day** 20d ann vol and/or ATR(14)/close (no lookahead).

Example::

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

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,
    list_parquet_symbols,
    load_equity_panels,
)
from RenTech.strategy_stack.cm_intraday_dip import (
    pdl_touch_short_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("--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("--out-prefix", type=Path, default=LOGS / "cm_intraday_pdl_short_vol_sweep_sp100")
    args = ap.parse_args()

    have = set(list_parquet_symbols(args.data_dir))
    sp100 = [t.upper() for t in get_sp100_tickers() if t.upper() in have][: int(args.max_tickers)]
    warmup = (pd.Timestamp(args.start) - pd.Timedelta(days=500)).strftime("%Y-%m-%d")
    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)

    syms = ["SPY"] + sp100
    intra, daily = load_equity_panels(syms, data_dir=args.data_dir, start=warmup, end=args.end)
    spy = intra.pop("SPY")
    daily.pop("SPY", None)
    n_book = len(intra)

    base = replace(pdl_touch_short_config(), max_concurrent=n_book, slippage_bps=3.0)
    variants = [("baseline", base)]
    for thr in [18, 20, 22, 25, 28, 30]:
        variants.append((f"vol20>={thr}%", replace(base, min_prior_vol_ann_pct=float(thr))))
    for thr in [1.5, 1.75, 2.0, 2.25, 2.5]:
        variants.append((f"atr_pct>={thr}%", replace(base, min_prior_atr_pct=float(thr))))
    for v, a in [(22, 1.5), (25, 1.5), (25, 2.0), (28, 2.0), (30, 2.0)]:
        variants.append(
            (
                f"vol>={v}% & atr>={a}%",
                replace(base, min_prior_vol_ann_pct=float(v), min_prior_atr_pct=float(a)),
            )
        )

    rows = []
    for label, cfg in variants:
        t0 = time.perf_counter()
        port, tr = run_backtest(intra, daily, cfg=cfg, spy_intra=spy, return_start=ret_start)
        port = port.loc[:end]
        pm = portfolio_metrics(port)
        ts = trade_stats(tr)
        row = {
            "label": label,
            "elapsed_sec": round(time.perf_counter() - t0, 1),
            **pm,
            **{f"trade_{k}": v for k, v in ts.items()},
            "min_prior_vol_ann_pct": cfg.min_prior_vol_ann_pct,
            "min_prior_atr_pct": cfg.min_prior_atr_pct,
        }
        rows.append(row)
        print(
            f"{label:22s} ret {pm.get('total_return_pct', 0):+7.1f}%  "
            f"Sh {pm.get('sharpe', 0):5.2f}  DD {pm.get('max_dd_pct', 0):6.1f}%  "
            f"n={ts.get('n_trades', 0):5d}  win={ts.get('win_rate_pct', 0):4.0f}%  "
            f"avg={ts.get('avg_pnl_pct', 0):+.3f}%",
            flush=True,
        )

    df = pd.DataFrame(rows).sort_values("sharpe", ascending=False)
    out = args.out_prefix.expanduser().resolve()
    df.to_csv(f"{out}.csv", index=False)
    cmd = (
        f"cd {_REPO} && PYTHONUNBUFFERED=1 .venv/bin/python "
        f"RenTech/strategy_stack/run_cm_intraday_pdl_short_vol_sweep.py "
        f"--start {args.start} --end {args.end} --out-prefix {out}"
    )
    Path(f"{out}_meta.json").write_text(json.dumps({"command": cmd, "rows": rows}, indent=2, default=str) + "\n")
    print(f"\nWrote {out}.csv", flush=True)


if __name__ == "__main__":
    main()
