#!/usr/bin/env python3
"""
Sweep five confirm_entry_4b enhancement ideas vs the current best config.

Ideas tested (all relative to confirm_entry_4b @ top-10 unless noted):
  1. require_above_vwap at confirm
  2. confirm_lag_bars=5 (later entry)
  3. top_n=5 (fewer names)
  4. require_price_rising_confirm (close[j2] > close[j1])
  5. min_rank_score_pct_gap=0.05 (sharp top-N cutoff)

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_confirm_ideas_sweep.py \\
        --start 2020-01-02 --end 2026-06-26
"""

from __future__ import annotations

import argparse
import json
import sys
from dataclasses import asdict, 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.ma_slope_intraday_enhanced import (
    EnhancedIntradayConfig,
    EnhancedIntradayEngine,
    baseline_enhanced_config,
    metrics_daily,
)
from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

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


def _confirm_base() -> EnhancedIntradayConfig:
    return replace(baseline_enhanced_config(), hold_mode="confirm_entry", confirm_lag_bars=4)


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__.split("\n\n")[0])
    ap.add_argument("--start", default="2020-01-02")
    ap.add_argument("--end", default="2026-06-26")
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_confirm_ideas_sweep")
    args = ap.parse_args()

    syms = list_parquet_symbols(DEFAULT_ALPACA_RTH_DIR)[: int(args.max_tickers)]
    if "SPY" not in syms:
        syms.insert(0, "SPY")
    sector_map = dict(
        zip(
            load_sp500_sectors()["ticker"].astype(str).str.upper(),
            load_sp500_sectors()["sector"].astype(str),
        )
    )

    print("Loading panels …", flush=True)
    intra, _ = load_equity_panels(
        syms,
        start="2019-11-01",
        end=str(args.end),
        warmup_sessions=15,
    )
    ret_start = pd.Timestamp(args.start)
    end = pd.Timestamp(args.end)

    base = _confirm_base()
    variants: list[tuple[str, EnhancedIntradayConfig, int]] = [
        ("confirm_4b_ref", base, 10),
        ("idea1_vwap", replace(base, require_above_vwap=True), 10),
        ("idea2_lag5", replace(base, confirm_lag_bars=5), 10),
        ("idea3_top5", base, 5),
        ("idea4_rising", replace(base, require_price_rising_confirm=True), 10),
        ("idea5_rank_gap5pct", replace(base, min_rank_score_pct_gap=0.05), 10),
    ]

    print("Building score panels (once) …", flush=True)
    eng0 = EnhancedIntradayEngine(config=base)
    panels = eng0.build_panels(intra, sector_map)

    rows = []
    ref_ret = None
    for name, cfg, top_n in variants:
        print(f"  {name} (top_n={top_n}) …", flush=True)
        eng = EnhancedIntradayEngine(config=cfg)
        port_r = eng.run(intra, top_n, panels=panels, return_start=ret_start)
        port_r = port_r.loc[port_r.index <= end]
        m = metrics_daily(port_r)
        row = {
            "variant": name,
            "top_n": top_n,
            **m,
            **{f"cfg_{k}": v for k, v in asdict(cfg).items() if k in (
                "require_above_vwap", "confirm_lag_bars", "require_price_rising_confirm",
                "min_rank_score_pct_gap",
            )},
        }
        if name == "confirm_4b_ref":
            ref_ret = float(m.get("total_return_pct", 0))
        rows.append(row)
        print(
            f"    ret {m.get('total_return_pct', float('nan')):+7.1f}%  "
            f"Sharpe {m.get('sharpe', float('nan')):5.2f}  DD {m.get('max_dd_pct', float('nan')):6.1f}%",
            flush=True,
        )

    df = pd.DataFrame(rows)
    if ref_ret is not None:
        df["return_vs_confirm_4b_pct"] = df["total_return_pct"] - ref_ret
    df = df.sort_values("total_return_pct", ascending=False)

    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_loaded": len(intra),
        "reference": "confirm_4b_ref",
        "reference_return_pct": ref_ret,
        "ranked": df.to_dict(orient="records"),
    }
    Path(f"{out}_meta.json").write_text(json.dumps(meta, indent=2) + "\n")

    print(f"\nReference confirm_4b: {ref_ret:+.1f}%", flush=True)
    print("\nRanked by return:", flush=True)
    for _, r in df.iterrows():
        delta = r.get("return_vs_confirm_4b_pct", 0)
        print(
            f"  {r['variant']:22s}  {r['total_return_pct']:+7.1f}%  "
            f"Δ {delta:+6.1f}pp  Sharpe {r['sharpe']:.2f}  DD {r['max_dd_pct']:.1f}%",
            flush=True,
        )
    print(f"\nWrote {csv_path}", flush=True)


if __name__ == "__main__":
    main()
