#!/usr/bin/env python3
"""
Baseline vs tweaked CM intraday dip sweep (+ single-name sanity checks).

Tweaked preset: SPY>VWAP + prior VIX<20, confirm hammer/vol/higher-low,
2.5×ATR stretch, 1% SPY underperf (C), daily ATR exits 1.0/0.75.

Example::

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

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,
    load_equity_panels,
    list_parquet_symbols,
)
from RenTech.strategy_stack.cm_intraday_dip import (
    baseline_config,
    portfolio_metrics,
    run_backtest,
    trade_stats,
    tweaked_config,
)
from universe_scanner import get_sp100_tickers

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


def _resolve_sp100(data_dir: Path, max_tickers: int) -> list[str]:
    have = set(list_parquet_symbols(data_dir))
    syms = [t.upper() for t in get_sp100_tickers() if t.upper() in have]
    if "SPY" not in syms and "SPY" in have:
        syms.insert(0, "SPY")
    return syms[: int(max_tickers)]


def _run_cfg(
    label: str,
    cfg,
    intra: dict,
    daily: dict,
    spy_intra: pd.DataFrame | None,
    *,
    ret_start: pd.Timestamp,
    end: pd.Timestamp,
) -> dict:
    t0 = time.perf_counter()
    port, trades = run_backtest(intra, daily, cfg=cfg, spy_intra=spy_intra, return_start=ret_start)
    port = port.loc[port.index <= end]
    pm = portfolio_metrics(port)
    ts = trade_stats(trades)
    return {
        "label": label,
        "variant": cfg.variant,
        "preset": "tweaked" if cfg.require_spy_above_vwap else "baseline",
        "elapsed_sec": round(time.perf_counter() - t0, 1),
        **pm,
        **{f"trade_{k}": v for k, v in ts.items()},
    }


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("--top-n", 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_dip_tweak_sweep")
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    symbols = _resolve_sp100(data_dir, args.max_tickers)
    warmup = (pd.Timestamp(args.start) - pd.Timedelta(days=400)).strftime("%Y-%m-%d")
    print(f"Loading {len(symbols)} SP100 names …", flush=True)
    intra_all, daily_all = load_equity_panels(
        symbols, data_dir=data_dir, start=warmup, end=args.end, verbose=True
    )
    spy_all = intra_all.pop("SPY", None)
    daily_all.pop("SPY", None)

    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)
    rows: list[dict] = []

    for variant in ("vwap_reclaim", "pdl_touch", "relative_washout"):
        for preset_name, factory in (("baseline", baseline_config), ("tweaked", tweaked_config)):
            cfg = replace(
                factory(variant),
                top_n=int(args.top_n),
                max_concurrent=int(args.top_n),
                slippage_bps=float(args.slippage_bps),
            )
            label = f"{variant}_{preset_name}_sp100"
            print(f"\n--- {label} ---", flush=True)
            row = _run_cfg(label, cfg, intra_all, daily_all, spy_all, ret_start=ret_start, end=end)
            rows.append(row)
            print(
                f"  ret {row.get('total_return_pct', float('nan')):+7.1f}%  "
                f"Sharpe {row.get('sharpe', float('nan')):5.2f}  "
                f"DD {row.get('max_dd_pct', float('nan')):6.1f}%  "
                f"trades={row.get('trade_n_trades', 0)}  "
                f"win={row.get('trade_win_rate_pct', float('nan')):.1f}%  "
                f"avg={row.get('trade_avg_pnl_pct', float('nan')):+.3f}%",
                flush=True,
            )

    singles = ("AAPL", "MSFT", "NVDA")
    for sym in singles:
        if sym not in intra_all:
            continue
        intra1 = {sym: intra_all[sym]}
        daily1 = {sym: daily_all[sym]}
        for variant in ("vwap_reclaim", "pdl_touch"):
            cfg = replace(
                tweaked_config(variant),
                max_concurrent=1,
                top_n=1,
                slippage_bps=float(args.slippage_bps),
            )
            label = f"{variant}_tweaked_{sym}"
            print(f"\n--- {label} ---", flush=True)
            row = _run_cfg(label, cfg, intra1, daily1, spy_all, ret_start=ret_start, end=end)
            rows.append(row)
            print(
                f"  ret {row.get('total_return_pct', float('nan')):+7.1f}%  "
                f"trades={row.get('trade_n_trades', 0)}  "
                f"win={row.get('trade_win_rate_pct', float('nan')):.1f}%  "
                f"avg={row.get('trade_avg_pnl_pct', float('nan')):+.3f}%",
                flush=True,
            )

    df = pd.DataFrame(rows)
    out = args.out_prefix.expanduser().resolve()
    csv_path = Path(f"{out}.csv")
    df.to_csv(csv_path, index=False)
    meta = {
        "start": args.start,
        "end": args.end,
        "n_sp100": len(intra_all),
        "slippage_bps": args.slippage_bps,
        "rows": rows,
    }
    Path(f"{out}_meta.json").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_tweak_sweep.py "
        f"--start {args.start} --end {args.end} --max-tickers {args.max_tickers} "
        f"--out-prefix {out}"
    )
    Path(f"{out}_metrics.txt").write_text(
        "\n".join([f"command: {cmd}", "", df.to_string(index=False), ""]) + "\n"
    )
    print(f"\nWrote {csv_path}", flush=True)


if __name__ == "__main__":
    main()
