#!/usr/bin/env python3
"""
**Same-day** MA slope rotation on Alpaca 5-minute bars.

* EMA periods are **5m bar counts** (default 10 / 50 ≈ 50 min / 4.2 hr).
* Prior sessions' 5m closes feed the MA (continuous series).
* **Enter** minutes **30–90** after 09:30 ET (bars 5–17); **hold to MOC** (same day, no overnight).
* Separate from the daily swing MA slope sleeves in the stock-only fund.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_ma_slope_alpaca_intraday_daytrade.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 dataclasses import asdict
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_cross_sectional import (
    MaSlopeCrossSectional,
    MaSlopeCrossSectionalConfig,
)
from RenTech.strategy_stack.ma_slope_intraday_daytrade import (
    DEFAULT_ENTRY_BAR,
    DEFAULT_SESSION_ENTRY_BAR_MAX,
    DEFAULT_SESSION_ENTRY_BAR_MIN,
    DEFAULT_SESSION_MAX_HOLD_BAR,
    MaSlopeIntradayDayTrade,
    MaSlopeIntradayDayTradeConfig,
)

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


def _metrics(r: pd.Series, *, capital: float = 100_000.0, ann_factor: float = 252.0) -> dict:
    r = r.astype(np.float64).dropna()
    if len(r) < 2:
        return {}
    eq = capital * (1.0 + r).cumprod()
    n = len(r)
    years = n / ann_factor
    end = float(eq.iloc[-1])
    tot = end / capital - 1.0
    cagr = (end / capital) ** (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(ann_factor)) if sd > 1e-12 else float("nan")
    return {
        "n": int(n),
        "total_return_pct": float(tot * 100.0),
        "cagr_pct": float(cagr * 100.0),
        "max_dd_pct": float(dd * 100.0),
        "sharpe": sharpe,
        "daily_win_rate_pct": float((r > 0).mean() * 100.0),
        "end_equity_usd": end,
    }


def _yearly(r: pd.Series) -> list[dict]:
    rows = []
    for yr, g in r.groupby(r.index.year):
        eq = (1.0 + g).cumprod()
        rows.append(
            {
                "year": int(yr),
                "return_pct": float((eq.iloc[-1] - 1.0) * 100.0),
                "max_dd_pct": float((eq / eq.cummax() - 1.0).min() * 100.0),
            }
        )
    return rows


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("--capital", type=float, default=100_000.0)
    ap.add_argument("--max-tickers", type=int, default=0)
    ap.add_argument("--symbols", default="")
    ap.add_argument("--warmup-sessions", type=int, default=15, help="Prior sessions for MA seed")
    ap.add_argument("--fast-period", type=int, default=10)
    ap.add_argument("--slow-period", type=int, default=50)
    ap.add_argument("--fast-lookback", type=int, default=10)
    ap.add_argument("--slow-lookback", type=int, default=5)
    ap.add_argument(
        "--hold-mode",
        default="once_per_session",
        choices=("once_per_session", "each_bar", "every_n_bars"),
        help="once_per_session: rank once after entry-bar, hold to close (default)",
    )
    ap.add_argument("--entry-bar", type=int, default=DEFAULT_ENTRY_BAR)
    ap.add_argument(
        "--session-entry-bar-min",
        type=int,
        default=DEFAULT_SESSION_ENTRY_BAR_MIN,
        help="Earliest entry bar in session (5 = ~10:00 ET exec)",
    )
    ap.add_argument(
        "--session-entry-bar-max",
        type=int,
        default=DEFAULT_SESSION_ENTRY_BAR_MAX,
        help="Latest entry signal bar (17 = ~11:00 ET exec)",
    )
    ap.add_argument(
        "--flat-at-90min",
        action="store_true",
        help="Flat at ~11:00 instead of holding to MOC (sets hold_to_moc=False)",
    )
    ap.add_argument(
        "--session-max-hold-bar",
        type=int,
        default=DEFAULT_SESSION_MAX_HOLD_BAR,
        help="With --flat-at-90min: last bar with exposure (17 = ~11:00)",
    )
    ap.add_argument("--rebalance-every-bars", type=int, default=6, help="With every_n_bars hold mode")
    ap.add_argument("--skip-bars", type=int, default=0, help="Extra skip before first rebalance (each_bar mode)")
    ap.add_argument(
        "--stop-mode",
        default="none",
        choices=("none", "atr_trail", "pct_trail", "fixed_pct_entry"),
    )
    ap.add_argument("--atr-period", type=int, default=14)
    ap.add_argument("--atr-multiplier", type=float, default=2.0)
    ap.add_argument("--pct-trail-stop", type=float, default=0.02, help="e.g. 0.02 = 2%% from session high")
    ap.add_argument("--fixed-stop-pct", type=float, default=0.015, help="e.g. 0.015 = 1.5%% from entry")
    ap.add_argument("--out-prefix", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    data_dir = args.data_dir.expanduser().resolve()
    if not data_dir.is_dir():
        raise SystemExit(f"Missing data dir: {data_dir}")

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

    print(f"Loading {len(symbols)} symbols (warmup {args.warmup_sessions} sessions) …", flush=True)
    intra_dict, daily_dict = load_equity_panels(
        symbols,
        data_dir=data_dir,
        bar_minutes=5,
        start=args.start,
        end=args.end,
        warmup_sessions=int(args.warmup_sessions),
    )
    if len(intra_dict) < max(1, int(args.top_n)):
        raise SystemExit(f"Only {len(intra_dict)} symbols loaded")

    ret_start = pd.Timestamp(args.start)
    print("\n=== 5m same-day MA slope (enter 30–90 min, hold to MOC) ===", flush=True)
    day_cfg = MaSlopeIntradayDayTradeConfig(
        fast_period=int(args.fast_period),
        slow_period=int(args.slow_period),
        fast_lookback=int(args.fast_lookback),
        slow_lookback=int(args.slow_lookback),
        skip_bars_per_session=int(args.skip_bars),
        hold_mode=args.hold_mode,  # type: ignore[arg-type]
        entry_bar=int(args.entry_bar),
        session_entry_bar_min=int(args.session_entry_bar_min),
        session_entry_bar_max=int(args.session_entry_bar_max),
        hold_to_moc=not bool(args.flat_at_90min),
        session_max_hold_bar=int(args.session_max_hold_bar),
        rebalance_every_bars=int(args.rebalance_every_bars),
        stop_mode=args.stop_mode,  # type: ignore[arg-type]
        atr_period=int(args.atr_period),
        atr_multiplier=float(args.atr_multiplier),
        pct_trail_stop=float(args.pct_trail_stop),
        fixed_stop_pct=float(args.fixed_stop_pct),
    )
    r_intra = MaSlopeIntradayDayTrade(config=day_cfg).generate_returns(
        intra_dict,
        top_n=int(args.top_n),
        return_start=ret_start,
    )
    if args.end.strip():
        r_intra = r_intra.loc[r_intra.index <= pd.Timestamp(args.end)]

    r_daily_sess = compound_intraday_to_daily(r_intra)
    r_daily_sess = r_daily_sess.loc[r_daily_sess.index >= ret_start.normalize()]
    if args.end.strip():
        r_daily_sess = r_daily_sess.loc[r_daily_sess.index <= pd.Timestamp(args.end).normalize()]

    print("\n=== Daily-bar monthly top-N (swing reference, same tickers) ===", flush=True)
    swing_cfg = MaSlopeCrossSectionalConfig(
        fast_period=int(args.fast_period),
        slow_period=int(args.slow_period),
        fast_lookback=int(args.fast_lookback),
        slow_lookback=int(args.slow_lookback),
        stop_mode="atr_trail",
        atr_multiplier=2.0,
    )
    r_swing = MaSlopeCrossSectional(config=swing_cfg).generate_returns(
        daily_dict, top_n=int(args.top_n)
    )
    r_swing = r_swing.loc[r_swing.index >= ret_start]
    if args.end.strip():
        r_swing = r_swing.loc[r_swing.index <= pd.Timestamp(args.end)]

    m_day = _metrics(r_daily_sess, capital=args.capital, ann_factor=252.0)
    m_swing = _metrics(r_swing, capital=args.capital, ann_factor=252.0)

    spy_ret = pd.Series(dtype=float)
    if "SPY" in daily_dict:
        spy_ret = daily_dict["SPY"]["ret"].astype(float)
        spy_ret.index = pd.to_datetime(spy_ret.index).tz_localize(None)
        spy_ret = spy_ret.loc[r_swing.index.intersection(spy_ret.index)]
    m_spy = _metrics(spy_ret, capital=args.capital) if len(spy_ret) else {}

    print("\n--- Session-compounded daily returns ---")
    for label, m in [
        ("5m day-trade", m_day),
        ("Daily swing top-N", m_swing),
        ("SPY B&H", m_spy),
    ]:
        if not m:
            continue
        print(
            f"  {label:22s}  ret {m['total_return_pct']:+7.1f}%  "
            f"CAGR {m['cagr_pct']:+6.1f}%  Sharpe {m['sharpe']:5.2f}  "
            f"MaxDD {m['max_dd_pct']:6.1f}%  win% {m.get('daily_win_rate_pct', 0):.1f}"
        )

    prefix = args.out_prefix.expanduser().resolve()
    prefix.parent.mkdir(parents=True, exist_ok=True)
    hold_tag = "moc" if day_cfg.hold_to_moc else f"flat{args.session_max_hold_bar}"
    stop_tag = day_cfg.stop_mode if day_cfg.stop_mode == "none" else (
        f"atr{day_cfg.atr_multiplier:g}x" if day_cfg.stop_mode == "atr_trail"
        else f"pct{int(day_cfg.pct_trail_stop*100)}"
        if day_cfg.stop_mode == "pct_trail"
        else f"fix{int(day_cfg.fixed_stop_pct*1000)/10:g}pct"
    )
    slug = f"top{int(args.top_n)}_n{len(intra_dict)}_win{args.session_entry_bar_min}_{args.session_entry_bar_max}_{hold_tag}_{stop_tag}"
    meta = {
        "window": {"start": args.start, "end": args.end},
        "n_symbols": len(intra_dict),
        "daytrade_config": asdict(day_cfg),
        "metrics_session_daily": {
            "intraday_daytrade": m_day,
            "daily_swing_reference": m_swing,
            "spy": m_spy,
        },
        "yearly_daytrade": _yearly(r_daily_sess),
        "yearly_swing": _yearly(r_swing),
    }
    meta_path = Path(f"{prefix}_{slug}_meta.json")
    meta_path.write_text(json.dumps(meta, indent=2) + "\n")
    pd.DataFrame({"datetime": r_intra.index, "bar_ret": r_intra.values}).to_csv(
        f"{prefix}_{slug}_bar_returns.csv", index=False
    )
    pd.DataFrame({"date": r_daily_sess.index, "daily_ret": r_daily_sess.values}).to_csv(
        f"{prefix}_{slug}_session_daily.csv", index=False
    )
    print(f"\nWrote {meta_path}")


if __name__ == "__main__":
    main()
