#!/usr/bin/env python3
"""
Sweep **entry bar** (time of day) for Alpaca 5m same-day MA slope top-N.

Signal at session bar ``entry_bar`` (0 = first 5m bar at 09:30); position from next bar.
Flat at session close. Prior sessions feed the MA.

Example::

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

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import numpy as np
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.ma_slope_intraday_daytrade import (
    MaSlopeIntradayDayTrade,
    MaSlopeIntradayDayTradeConfig,
)
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames

LOGS = _REPO / "RenTech" / "data" / "logs"
RTH_BARS = 78  # 09:30–16:00 on 5m grid


def entry_bar_to_exec_time(entry_bar: int) -> str:
    """First execution bar open (US/Eastern), session opens 09:30."""
    mins = 9 * 60 + 30 + (int(entry_bar) + 1) * 5
    h, m = divmod(mins, 60)
    return f"{h:02d}:{m:02d}"


def _session_metrics(r_daily: pd.Series) -> dict:
    r = r_daily.astype(np.float64).dropna()
    if len(r) < 2:
        return {}
    eq = (1.0 + r).cumprod()
    years = len(r) / 252.0
    tot = float(eq.iloc[-1] - 1.0)
    cagr = float(eq.iloc[-1] ** (1.0 / years) - 1.0) if years > 0 else float("nan")
    dd = float((eq / eq.cummax() - 1.0).min())
    sd = float(r.std(ddof=1))
    sharpe = float(r.mean() / sd * np.sqrt(252.0)) if sd > 1e-12 else float("nan")
    return {
        "total_return_pct": tot * 100.0,
        "cagr_pct": cagr * 100.0,
        "max_dd_pct": dd * 100.0,
        "sharpe": sharpe,
        "daily_win_rate_pct": float((r > 0).mean() * 100.0),
        "n_sessions": int(len(r)),
    }


def _returns_for_entry(
    eng: MaSlopeIntradayDayTrade,
    equity_dict: dict,
    score_df: pd.DataFrame,
    top_n: int,
    entry_bar: int,
    return_start: pd.Timestamp,
) -> pd.Series:
    cfg = eng.config
    eng_cfg = MaSlopeIntradayDayTradeConfig(
        **{
            **cfg.__dict__,
            "entry_bar": int(entry_bar),
            "hold_mode": "once_per_session",
        }
    )
    eng2 = MaSlopeIntradayDayTrade(config=eng_cfg)

    ret_pan = []
    for t, df in sorted(equity_dict.items()):
        idx = pd.to_datetime(df.index).tz_localize(None)
        dfx = df.copy()
        dfx.index = idx
        ret_pan.append(dfx["ret"].astype(np.float64).rename(t))
    ret_df, _ = _align_panel_frames(ret_pan)
    master = ret_df.index.sort_values()
    ret_df = ret_df.reindex(master).fillna(0.0)
    score_df = score_df.reindex(master).reindex(columns=ret_df.columns)

    target_w = eng2.target_weights(score_df, top_n)
    exec_w = target_w.shift(1).fillna(0.0)
    first = eng2._first_bar_mask(master)
    exec_w.loc[first] = 0.0
    w = exec_w.to_numpy(dtype=np.float64)
    r = ret_df.to_numpy(dtype=np.float64)
    port_r = pd.Series((w * r).sum(axis=1), index=master)
    return port_r.loc[port_r.index >= return_start]


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("--symbols", default="")
    ap.add_argument("--warmup-sessions", type=int, default=15)
    ap.add_argument(
        "--entry-bars",
        default="",
        help="Comma-separated bar indices (default: 0,2,4,...,72 step 2)",
    )
    ap.add_argument("--out-prefix", type=Path, default=LOGS / "ma_slope_alpaca_entry_time_sweep")
    args = ap.parse_args()

    if args.entry_bars.strip():
        entry_bars = [int(x) for x in args.entry_bars.split(",") if x.strip()]
    else:
        from RenTech.strategy_stack.ma_slope_intraday_daytrade import (
            DEFAULT_SESSION_ENTRY_BAR_MAX,
            DEFAULT_SESSION_ENTRY_BAR_MIN,
        )
        entry_bars = list(range(DEFAULT_SESSION_ENTRY_BAR_MIN, DEFAULT_SESSION_ENTRY_BAR_MAX + 1, 2))

    data_dir = args.data_dir.expanduser().resolve()
    if args.symbols.strip():
        symbols = [s.strip().upper() for s in args.symbols.split(",") if s.strip()]
    else:
        symbols = list_parquet_symbols(data_dir)[: int(args.max_tickers)]

    print(f"Loading {len(symbols)} symbols …", flush=True)
    intra_dict, _ = load_equity_panels(
        symbols,
        data_dir=data_dir,
        bar_minutes=5,
        start=args.start,
        end=args.end,
        warmup_sessions=int(args.warmup_sessions),
    )
    ret_start = pd.Timestamp(args.start)
    base_cfg = MaSlopeIntradayDayTradeConfig(hold_mode="once_per_session")
    eng = MaSlopeIntradayDayTrade(config=base_cfg)
    print("Building slope score panel (once) …", flush=True)
    score_df = eng.build_score_panel(intra_dict)

    rows: list[dict] = []
    for eb in entry_bars:
        port_r = _returns_for_entry(
            eng, intra_dict, score_df, int(args.top_n), eb, ret_start
        )
        if args.end.strip():
            port_r = port_r.loc[port_r.index <= pd.Timestamp(args.end)]
        r_sess = compound_intraday_to_daily(port_r)
        r_sess = r_sess.loc[r_sess.index >= ret_start.normalize()]
        m = _session_metrics(r_sess)
        mins_in = 30 + eb * 5
        bucket = "first_30m" if eb < 6 else ("10_11" if eb < 12 else ("11_12" if eb < 18 else ("12_14" if eb < 30 else "after_14")))
        row = {
            "entry_bar": eb,
            "signal_bar_end_et": entry_bar_to_exec_time(eb - 1) if eb > 0 else "09:35",
            "first_exec_et": entry_bar_to_exec_time(eb),
            "minutes_from_open": mins_in,
            "bucket": bucket,
            **m,
        }
        rows.append(row)
        print(
            f"  bar {eb:2d}  exec ~{row['first_exec_et']}  "
            f"ret {m.get('total_return_pct', float('nan')):+7.1f}%  "
            f"Sharpe {m.get('sharpe', float('nan')):5.2f}  "
            f"DD {m.get('max_dd_pct', float('nan')):6.1f}%",
            flush=True,
        )

    df = pd.DataFrame(rows).sort_values("entry_bar")
    best_sh = df.loc[df["sharpe"].idxmax()] if df["sharpe"].notna().any() else None
    best_ret = df.loc[df["total_return_pct"].idxmax()] if len(df) else None

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    slug = f"top{int(args.top_n)}_n{len(intra_dict)}"
    csv_path = Path(f"{prefix}_{slug}.csv")
    df.to_csv(csv_path, index=False)

    summary = {
        "window": {"start": args.start, "end": args.end},
        "n_symbols": len(intra_dict),
        "top_n": int(args.top_n),
        "best_sharpe": best_sh.to_dict() if best_sh is not None else None,
        "best_return": best_ret.to_dict() if best_ret is not None else None,
        "bucket_means": df.groupby("bucket")[["total_return_pct", "sharpe", "max_dd_pct"]].mean().to_dict(),
        "first_30m_vs_after": {
            "first_30m_mean_sharpe": float(df.loc[df["entry_bar"] < 6, "sharpe"].mean()),
            "after_30m_mean_sharpe": float(df.loc[df["entry_bar"] >= 6, "sharpe"].mean()),
            "first_30m_mean_return": float(df.loc[df["entry_bar"] < 6, "total_return_pct"].mean()),
            "after_30m_mean_return": float(df.loc[df["entry_bar"] >= 6, "total_return_pct"].mean()),
        },
    }
    json_path = Path(f"{prefix}_{slug}_meta.json")
    json_path.write_text(json.dumps(summary, indent=2) + "\n")

    print("\n--- Best by Sharpe ---")
    if best_sh is not None:
        print(
            f"  bar {int(best_sh['entry_bar'])}  exec ~{best_sh['first_exec_et']}  "
            f"ret {best_sh['total_return_pct']:+.1f}%  Sharpe {best_sh['sharpe']:.2f}"
        )
    print(f"\nWrote {csv_path}")
    print(f"Wrote {json_path}")


if __name__ == "__main__":
    main()
