#!/usr/bin/env python3
"""
Sweep intraday MA-slope **enhancements** vs baseline (top-10, enter 30–90, hold MOC).

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_alpaca_intraday_enhancement_sweep.py \\
        --start 2020-01-02 --end 2024-12-31 --top-n 10
"""

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,
    liquidity_filter_symbols,
    load_sp500_parquet_tickers,
    metrics_daily,
)
from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

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


def _ensure_spy(symbols: list[str]) -> list[str]:
    out = list(symbols)
    if "SPY" not in out:
        out.insert(0, "SPY")
    return out


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("--start", default="2020-01-02")
    ap.add_argument("--end", default="2024-12-31")
    ap.add_argument("--top-n", type=int, default=10)
    ap.add_argument("--max-tickers", type=int, default=500)
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_intraday_enhancement_sweep")
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    all_syms = list_parquet_symbols(data_dir)[: int(args.max_tickers)]
    sp500_syms = load_sp500_parquet_tickers(data_dir)
    sec_df = load_sp500_sectors()
    sector_map = dict(zip(sec_df["ticker"].astype(str).str.upper(), sec_df["sector"].astype(str)))

    print("Loading alphabetical universe …", flush=True)
    intra_all, daily_all = load_equity_panels(
        _ensure_spy(all_syms), data_dir=data_dir, start=args.start, end=args.end, warmup_sessions=15
    )
    print(f"Loading S&P 500 parquet universe ({len(sp500_syms)} names) …", flush=True)
    intra_sp, daily_sp = load_equity_panels(
        _ensure_spy(sp500_syms), data_dir=data_dir, start=args.start, end=args.end, warmup_sessions=15
    )

    liquid = liquidity_filter_symbols(daily_all, min_price=5.0, min_adv=5_000_000)
    intra_liq = {k: v for k, v in intra_all.items() if k in liquid}
    daily_liq = {k: v for k, v in daily_all.items() if k in liquid}

    ret_start = pd.Timestamp(args.start)
    top_n = int(args.top_n)

    variants: list[tuple[str, dict, dict[str, pd.DataFrame], dict[str, str]]] = []

    def add(name: str, cfg: EnhancedIntradayConfig, intra: dict, sec: dict[str, str] | None = None):
        variants.append((name, asdict(cfg), intra, sec or sector_map))

    base = baseline_enhanced_config()
    add("baseline_alpha500", base, intra_all)
    add("universe_sp500", base, intra_sp)
    add("liquidity_5m_adv", base, intra_liq)
    add("slippage_5bps", replace(base, slippage_bps=5.0), intra_all)
    add("slippage_10bps", replace(base, slippage_bps=10.0), intra_all)
    add("spy_dual_slope", replace(base, spy_filter="dual_slope"), intra_all)
    add("spy_vwap", replace(base, spy_filter="vwap"), intra_all)
    add("exit_slope_flip", replace(base, exit_mode="slope_flip"), intra_all)
    add("exit_price_cross", replace(base, exit_mode="price_cross"), intra_all)
    add("weight_score_prop", replace(base, weight_mode="score_prop"), intra_all)
    add("confirm_entry_4b", replace(base, hold_mode="confirm_entry", confirm_lag_bars=4), intra_all)
    add("confirm_entry_2b", replace(base, hold_mode="confirm_entry", confirm_lag_bars=2), intra_all)
    add("filter_vwap", replace(base, require_above_vwap=True), intra_all)
    add("filter_or_break", replace(base, require_or_break=True), intra_all)
    add("session_stop_1pct", replace(base, session_stop_pct=0.01), intra_all)
    add("session_stop_1p5pct", replace(base, session_stop_pct=0.015), intra_all)
    add("session_stop_2pct", replace(base, session_stop_pct=0.02), intra_all)
    add("spy_gross_scale", replace(base, spy_gross_scale=True), intra_all)
    add("sector_cap_2", replace(base, sector_cap=2), intra_all)
    add("sector_cap_3", replace(base, sector_cap=3), intra_all)
    # Combos that paired well in design
    add("combo_sp500_spy_slope", replace(base, spy_filter="dual_slope"), intra_sp)
    add("combo_sp500_score_weight", replace(base, weight_mode="score_prop"), intra_sp)
    add("combo_sp500_slope_exit", replace(base, exit_mode="slope_flip"), intra_sp)
    add("combo_sp500_vwap_or", replace(base, require_above_vwap=True, require_or_break=True), intra_sp)
    add(
        "combo_best_guess",
        replace(
            base,
            spy_filter="dual_slope",
            weight_mode="score_prop",
            exit_mode="slope_flip",
            require_above_vwap=True,
            slippage_bps=5.0,
        ),
        intra_sp,
    )
    add(
        "combo_sp500_confirm_spy",
        replace(base, hold_mode="confirm_entry", confirm_lag_bars=4, spy_filter="dual_slope"),
        intra_sp,
    )

    rows = []
    panels_cache: dict[int, object] = {}

    for name, cfg_dict, intra, sec in variants:
        cfg = EnhancedIntradayConfig(**cfg_dict)
        eng = EnhancedIntradayEngine(config=cfg)
        key = id(intra)
        if key not in panels_cache:
            print(f"  build panels n={len(intra)} …", flush=True)
            panels_cache[key] = eng.build_panels(intra, sec)
        panels = panels_cache[key]
        r = eng.run(intra, top_n, panels=panels, return_start=ret_start)
        if args.end.strip():
            r = r.loc[r.index <= pd.Timestamp(args.end)]
        m = metrics_daily(r)
        row = {"variant": name, "n_symbols": len(intra), **m}
        rows.append(row)
        print(
            f"  {name:28s}  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).sort_values("total_return_pct", ascending=False)
    baseline_ret = float(df.loc[df["variant"] == "baseline_alpha500", "total_return_pct"].iloc[0])
    df["return_vs_baseline_pct"] = df["total_return_pct"] - baseline_ret

    out = args.out_prefix.expanduser().resolve()
    csv_path = Path(f"{out}_top{top_n}.csv")
    df.to_csv(csv_path, index=False)
    meta = {
        "baseline_return_pct": baseline_ret,
        "best_return": df.iloc[0].to_dict(),
        "best_sharpe": df.sort_values("sharpe", ascending=False).iloc[0].to_dict(),
        "best_dd": df.sort_values("max_dd_pct", ascending=False).iloc[0].to_dict(),
    }
    Path(f"{out}_top{top_n}_meta.json").write_text(json.dumps(meta, indent=2) + "\n")

    print(f"\nBaseline: {baseline_ret:+.1f}%")
    print("\nTop 10 by return:")
    for _, r in df.head(10).iterrows():
        print(
            f"  {r['variant']:28s}  {r['total_return_pct']:+7.1f}%  "
            f"Δ {r['return_vs_baseline_pct']:+6.1f}pp  Sharpe {r['sharpe']:.2f}  DD {r['max_dd_pct']:.1f}%"
        )
    print(f"\nWrote {csv_path}")


if __name__ == "__main__":
    main()
