#!/usr/bin/env python3
"""
Intraday **ATR high breakout** long on 5-minute RTH bars.

Buy when price first crosses the session breakout stop:
  level = anchor + atr_mult × prior-day ATR

Anchors (CrackingMarkets-style default):
  max(session_open, SMA10_prior_close) + mult × ATR(14)_prior

Exit: session close (MOC) by default; optional ATR profit/stop.

Example::

    cd /Users/robzingale/trading_bot && PYTHONUNBUFFERED=1 .venv/bin/python \\
        RenTech/strategy_stack/run_cm_intraday_atr_breakout.py \\
        --symbols AAPL,MSFT,NVDA,AMZN,GOOGL \\
        --start 2022-01-03 --end 2025-12-31 \\
        --out-prefix RenTech/data/logs/cm_intraday_atr_breakout_mega5
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Literal

import numpy as np
import pandas as pd

from RenTech.strategy_stack.cm_intraday_dip import (
    RegimeContext,
    TradeRow,
    _can_exit,
    _exit_atr,
    _portfolio_open_count,
    _session_bar_index,
    _session_key,
    _slip,
    _wilder_atr,
    build_regime_context,
    portfolio_metrics,
    trade_stats,
)
from RenTech.strategy_stack.ma_slope_engine import compute_slope_rank_score
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames as _align

AnchorMode = Literal["open", "poi", "max_open_poi"]
ExitMode = Literal["moc", "atr"]


@dataclass(frozen=True)
class AtrBreakoutConfig:
    atr_period: int = 14
    poi_period: int = 10
    atr_mult: float = 1.0
    anchor: AnchorMode = "max_open_poi"
    use_open_filter: bool = True
    require_uptrend: bool = True
    min_session_bars: int = 1
    max_entry_bar: int = 48  # ~11:30 ET on 5m grid; 0 = no cap
    exit_mode: ExitMode = "moc"
    profit_atr_mult: float = 1.0
    stop_atr_mult: float = 0.75
    use_daily_atr_exits: bool = True
    min_bars_held_before_exit: int = 1
    slippage_bps: float = 3.0
    max_concurrent: int = 10
    require_spy_above_vwap: bool = False
    vix_max_prior: float = 0.0
    sma_trend: int = 200
    require_close_above_level: bool = False
    failed_break_exit_bars: int = 0
    min_breakout_volume_mult: float = 0.0
    max_entries_per_bar: int = 0
    require_nr7: bool = False
    nr7_lookback: int = 7
    slope_top_n: int = 0


def atr_breakout_config() -> AtrBreakoutConfig:
    """Default: session open / POI + 1× prior daily ATR, MOC exit."""
    return AtrBreakoutConfig()


def atr_breakout_early_config() -> AtrBreakoutConfig:
    """Early-session cap (~10:30 ET on 5m bars) + 1× prior daily ATR."""
    return AtrBreakoutConfig(max_entry_bar=12)


def atr_breakout_cm_config() -> AtrBreakoutConfig:
    """CrackingMarkets NASDAQ-style: POI + 2.2× ATR(20), open filter, MOC."""
    return AtrBreakoutConfig(
        atr_period=20,
        poi_period=10,
        atr_mult=2.2,
        anchor="poi",
        use_open_filter=True,
        exit_mode="moc",
    )


def build_slope_top_intraday_mask(
    master_index: pd.DatetimeIndex,
    daily: Dict[str, pd.DataFrame],
    symbols: list[str],
    *,
    top_n: int,
) -> pd.DataFrame:
    """Bool panel: symbol was in top-N daily slope ranks at prior session close."""
    score_rows: dict[str, pd.Series] = {}
    for sym, db in daily.items():
        if db is None or db.empty or "close" not in db.columns:
            continue
        c = db["close"].astype(np.float64)
        c.index = pd.to_datetime(c.index).tz_localize(None)
        scored = compute_slope_rank_score(c.sort_index())
        score_rows[sym] = scored["rank_score"]
    if not score_rows:
        return pd.DataFrame(False, index=master_index, columns=symbols)

    score_df = pd.DataFrame(score_rows).sort_index()
    top_daily = pd.DataFrame(False, index=score_df.index, columns=score_df.columns)
    for dt, row in score_df.iterrows():
        picks = row.dropna().nlargest(int(top_n))
        for sym in picks.index:
            top_daily.loc[dt, sym] = True
    top_daily = top_daily.shift(1).fillna(False).astype(bool)

    ses = pd.DatetimeIndex([pd.Timestamp(s).normalize() for s in _session_key(master_index)])
    out = pd.DataFrame(False, index=master_index, columns=symbols)
    for sym in symbols:
        if sym not in top_daily.columns:
            continue
        ser = top_daily[sym]
        out[sym] = ses.map(lambda d: bool(ser.get(d, False))).to_numpy()
    return out


def build_breakout_panels(
    intra: Dict[str, pd.DataFrame],
    daily: Dict[str, pd.DataFrame],
    *,
    cfg: AtrBreakoutConfig,
    slope_daily: Dict[str, pd.DataFrame] | None = None,
) -> dict[str, pd.DataFrame]:
    close_p, open_p, high_p, low_p, atr_p, datr_p = [], [], [], [], [], []
    uptrend_p, ses_open_p, poi_p, vol_p, nr7_p = [], [], [], [], []

    for sym in sorted(intra.keys()):
        db = daily.get(sym)
        ib = intra[sym]
        if db is None or db.empty:
            continue
        idx = pd.to_datetime(ib.index).tz_localize(None)
        c = ib["close"].astype(np.float64)
        o = ib["open"].astype(np.float64) if "open" in ib.columns else c.copy()
        h = ib["high"].astype(np.float64)
        l = ib["low"].astype(np.float64)
        vol = ib["volume"].astype(np.float64) if "volume" in ib.columns else pd.Series(1.0, index=idx)
        atr_i = _wilder_atr(h, l, c, cfg.atr_period)

        d_idx = pd.to_datetime(db.index).tz_localize(None)
        d_close = db["close"].astype(np.float64)
        d_high = db["high"].astype(np.float64)
        d_low = db["low"].astype(np.float64)
        sma = d_close.rolling(cfg.sma_trend, min_periods=cfg.sma_trend).mean()
        uptrend_d = (d_close.shift(1) > sma.shift(1)).astype(np.float64)
        poi_d = d_close.rolling(cfg.poi_period, min_periods=cfg.poi_period).mean().shift(1)
        d_atr_d = _wilder_atr(d_high, d_low, d_close, cfg.atr_period).shift(1)
        d_range = d_high - d_low
        nr7_d = (d_range == d_range.rolling(int(cfg.nr7_lookback), min_periods=int(cfg.nr7_lookback)).min()).shift(1)

        ses = _session_key(idx)
        ses_open = c.groupby(ses).transform("first")
        uptrend = ses.map(uptrend_d).astype(np.float64)
        poi = ses.map(poi_d).astype(np.float64)
        d_atr = ses.map(d_atr_d).astype(np.float64)
        nr7 = ses.map(nr7_d).astype(np.float64)

        close_p.append(c.rename(sym))
        open_p.append(o.rename(sym))
        high_p.append(h.rename(sym))
        low_p.append(l.rename(sym))
        atr_p.append(atr_i.rename(sym))
        datr_p.append(d_atr.rename(sym))
        ses_open_p.append(ses_open.rename(sym))
        uptrend_p.append(uptrend.rename(sym))
        poi_p.append(poi.rename(sym))
        vol_p.append(vol.rename(sym))
        nr7_p.append(nr7.rename(sym))

    if not close_p:
        raise ValueError("no symbols with data")

    close_df, cols = _align(close_p)
    open_df, _ = _align(open_p)
    high_df, _ = _align(high_p)
    low_df, _ = _align(low_p)
    atr_df, _ = _align(atr_p)
    datr_df, _ = _align(datr_p)
    ses_open_df, _ = _align(ses_open_p)
    uptrend_df, _ = _align(uptrend_p)
    poi_df, _ = _align(poi_p)
    vol_df, _ = _align(vol_p)
    nr7_df, _ = _align(nr7_p)

    out = {
        "close": close_df.reindex(columns=cols),
        "open": open_df.reindex(columns=cols),
        "high": high_df.reindex(columns=cols),
        "low": low_df.reindex(columns=cols),
        "atr": atr_df.reindex(columns=cols),
        "daily_atr": datr_df.reindex(columns=cols),
        "ses_open": ses_open_df.reindex(columns=cols),
        "uptrend": uptrend_df.reindex(columns=cols),
        "poi": poi_df.reindex(columns=cols),
        "volume": vol_df.reindex(columns=cols),
        "nr7": nr7_df.reindex(columns=cols),
        "columns": pd.Index(cols),
    }
    if int(cfg.slope_top_n) > 0:
        slope_src = slope_daily if slope_daily is not None else {s: daily[s] for s in cols if s in daily}
        out["slope_ok"] = build_slope_top_intraday_mask(
            close_df.index, slope_src, list(cols), top_n=int(cfg.slope_top_n)
        ).reindex(columns=cols)
    return out


def _session_entry_level(
    ses_open: float,
    poi: float,
    daily_atr: float,
    cfg: AtrBreakoutConfig,
) -> float:
    if not np.isfinite(daily_atr) or daily_atr <= 0:
        return float("nan")
    if cfg.anchor == "open":
        base = ses_open
    elif cfg.anchor == "poi":
        base = poi
    else:
        parts = [ses_open]
        if np.isfinite(poi):
            parts.append(poi)
        base = float(max(parts))
    if not np.isfinite(base):
        return float("nan")
    return base + float(cfg.atr_mult) * daily_atr


def simulate_atr_breakout(
    panels: dict[str, pd.DataFrame],
    cfg: AtrBreakoutConfig,
    regime: RegimeContext | None = None,
) -> tuple[list[TradeRow], np.ndarray]:
    close = panels["close"].to_numpy(dtype=np.float64)
    high = panels["high"].to_numpy(dtype=np.float64)
    low = panels["low"].to_numpy(dtype=np.float64)
    ses_open = panels["ses_open"].to_numpy(dtype=np.float64)
    poi = panels["poi"].to_numpy(dtype=np.float64)
    daily_atr = panels["daily_atr"].to_numpy(dtype=np.float64)
    uptrend = panels["uptrend"].to_numpy(dtype=np.float64)
    atr = panels["atr"].to_numpy(dtype=np.float64)
    volume = panels["volume"].to_numpy(dtype=np.float64)
    nr7 = panels["nr7"].to_numpy(dtype=np.float64)
    slope_ok = (
        panels["slope_ok"].to_numpy(dtype=bool)
        if "slope_ok" in panels
        else np.ones(close.shape, dtype=bool)
    )
    idx = panels["close"].index
    cols = list(panels["columns"])
    bar_in_ses = _session_bar_index(idx).to_numpy(dtype=np.int64)
    ses = _session_key(idx).to_numpy()

    t_rows, n_sym = close.shape[0], close.shape[1]
    port_r = np.zeros(t_rows, dtype=np.float64)
    trades: list[TradeRow] = []

    state = np.zeros(n_sym, dtype=np.int8)
    entry_px = np.full(n_sym, np.nan)
    entry_exit_atr = np.full(n_sym, np.nan)
    entry_bar = np.full(n_sym, -1, dtype=np.int32)
    entry_time = [None] * n_sym
    traded_session = np.full(n_sym, None, dtype=object)
    session_run_high = np.full(n_sym, np.nan)
    session_level = np.full(n_sym, np.nan)
    session_cum_vol = np.zeros(n_sym, dtype=np.float64)

    slip = float(cfg.slippage_bps)
    max_conc = int(cfg.max_concurrent)
    max_entry = int(cfg.max_entry_bar)
    max_per_bar = int(cfg.max_entries_per_bar)
    fail_bars = int(cfg.failed_break_exit_bars)
    vol_mult = float(cfg.min_breakout_volume_mult)

    prev_ses = None
    for d in range(t_rows):
        if prev_ses is None or ses[d] != prev_ses:
            state[:] = 0
            traded_session[:] = None
            entry_px[:] = np.nan
            entry_exit_atr[:] = np.nan
            entry_bar[:] = -1
            session_run_high[:] = np.nan
            session_level[:] = np.nan
            session_cum_vol[:] = 0.0
            for i in range(n_sym):
                if np.isfinite(ses_open[d, i]) and np.isfinite(daily_atr[d, i]):
                    session_level[i] = _session_entry_level(
                        ses_open[d, i], poi[d, i], daily_atr[d, i], cfg
                    )
            prev_ses = ses[d]

        for i in range(n_sym):
            if np.isfinite(high[d, i]):
                so = ses_open[d, i]
                session_run_high[i] = (
                    high[d, i]
                    if not np.isfinite(session_run_high[i])
                    else max(session_run_high[i], high[d, i])
                )
            if np.isfinite(volume[d, i]):
                session_cum_vol[i] += volume[d, i]

        # exits
        for i in range(n_sym):
            if state[i] != 1:
                continue
            c_d = close[d, i]
            if not np.isfinite(c_d):
                continue
            ep, ea = entry_px[i], entry_exit_atr[i]
            eb = int(entry_bar[i])
            reason = None
            if fail_bars > 0 and d > eb and (d - eb) <= fail_bars:
                lvl = session_level[i]
                if np.isfinite(lvl) and c_d < lvl:
                    reason = "failed_break"
            if cfg.exit_mode == "atr" and reason is None and _can_exit(d, eb, _breakout_exit_cfg(cfg)):
                prof_k, stop_k = float(cfg.profit_atr_mult), float(cfg.stop_atr_mult)
                if np.isfinite(ea) and c_d >= ep + prof_k * ea:
                    reason = "profit_atr"
                elif np.isfinite(ea) and c_d <= ep - stop_k * ea:
                    reason = "stop_atr"
            if d + 1 >= t_rows or ses[d + 1] != ses[d]:
                reason = reason or "session_close"
            if reason:
                xp = _slip(c_d, slip, "sell")
                trades.append(
                    TradeRow(
                        variant="atr_high_breakout",
                        symbol=cols[i],
                        session_date=str(pd.Timestamp(ses[d]).date()),
                        entry_time=str(entry_time[i]),
                        exit_time=str(idx[d]),
                        entry_price=float(ep),
                        exit_price=float(xp),
                        pnl_pct=float(xp / ep - 1.0) if ep > 0 else 0.0,
                        bars_held=int(d - entry_bar[i]),
                        exit_reason=reason,
                    )
                )
                state[i] = 0

        # entries: first bar that crosses breakout stop
        if bar_in_ses[d] >= int(cfg.min_session_bars) and (max_entry <= 0 or bar_in_ses[d] <= max_entry):
            cand: list[tuple[int, float]] = []
            for i in range(n_sym):
                if state[i] == 1 or traded_session[i] is not None:
                    continue
                if cfg.require_uptrend and uptrend[d, i] < 0.5:
                    continue
                if cfg.require_nr7 and nr7[d, i] < 0.5:
                    continue
                if not slope_ok[d, i]:
                    continue
                if not _regime_allows_breakout(d, regime, cfg):
                    continue
                lvl = session_level[i]
                h_d, l_d, c_d = high[d, i], low[d, i], close[d, i]
                if not (np.isfinite(lvl) and np.isfinite(h_d) and np.isfinite(l_d) and np.isfinite(c_d)):
                    continue
                if h_d < lvl:
                    continue
                if cfg.require_close_above_level and c_d < lvl:
                    continue
                if d > 0 and ses[d] == ses[d - 1] and np.isfinite(high[d - 1, i]) and high[d - 1, i] >= lvl:
                    continue
                if cfg.use_open_filter:
                    so = ses_open[d, i]
                    rh = session_run_high[i]
                    if not (np.isfinite(so) and np.isfinite(rh) and rh > so):
                        continue
                if vol_mult > 0:
                    bar_n = int(bar_in_ses[d])
                    if bar_n <= 0:
                        continue
                    avg_vol = session_cum_vol[i] / float(bar_n)
                    if not (np.isfinite(volume[d, i]) and avg_vol > 0 and volume[d, i] >= vol_mult * avg_vol):
                        continue
                ext = (h_d - lvl) / daily_atr[d, i] if np.isfinite(daily_atr[d, i]) and daily_atr[d, i] > 0 else 0.0
                cand.append((i, ext))
            cand.sort(key=lambda x: -x[1])
            slots = max(0, max_conc - _portfolio_open_count(state))
            if max_per_bar > 0:
                slots = min(slots, max_per_bar)
            for i, _ in cand[:slots]:
                lvl = session_level[i]
                o_d = float(panels["open"].to_numpy(dtype=np.float64)[d, i])
                raw_fill = lvl
                if np.isfinite(o_d) and o_d >= lvl:
                    raw_fill = o_d
                ep = _slip(raw_fill, slip, "buy")
                a_d = atr[d, i]
                state[i] = 1
                entry_px[i] = ep
                entry_exit_atr[i] = _exit_atr(a_d, daily_atr[d, i], _breakout_exit_cfg(cfg))
                entry_bar[i] = d
                entry_time[i] = str(idx[d])
                traded_session[i] = ses[d]

        held = state == 1
        if np.any(held):
            rets = np.zeros(n_sym, dtype=np.float64)
            for i in np.flatnonzero(held):
                fd = int(entry_bar[i])
                c_d = close[d, i]
                if fd == d:
                    rets[i] = c_d / entry_px[i] - 1.0 if entry_px[i] > 0 else 0.0
                elif d > 0 and np.isfinite(close[d - 1, i]) and close[d - 1, i] > 0:
                    rets[i] = c_d / close[d - 1, i] - 1.0
            active = np.isfinite(rets) & held
            if np.any(active):
                w = 1.0 / float(np.sum(active))
                port_r[d] = float(np.sum(rets[active] * w))

    return trades, port_r


def _breakout_exit_cfg(cfg: AtrBreakoutConfig):
    """Adapter for shared ``_can_exit`` (needs min_bars field)."""
    from RenTech.strategy_stack.cm_intraday_dip import CmIntradayConfig

    return CmIntradayConfig(min_bars_held_before_exit=cfg.min_bars_held_before_exit)


def _regime_allows_breakout(d: int, regime: RegimeContext | None, cfg: AtrBreakoutConfig) -> bool:
    if regime is None:
        return True
    if cfg.require_spy_above_vwap and not bool(regime.spy_above_vwap[d]):
        return False
    if float(cfg.vix_max_prior) > 0 and not bool(regime.vix_ok[d]):
        return False
    return True


def run_backtest(
    intra: Dict[str, pd.DataFrame],
    daily: Dict[str, pd.DataFrame],
    *,
    cfg: AtrBreakoutConfig,
    spy_intra: pd.DataFrame | None = None,
    slope_daily: Dict[str, pd.DataFrame] | None = None,
    return_start: pd.Timestamp | None = None,
) -> tuple[pd.Series, pd.DataFrame]:
    panels = build_breakout_panels(intra, daily, cfg=cfg, slope_daily=slope_daily)
    master = panels["close"].index
    regime = None
    if spy_intra is not None and (cfg.require_spy_above_vwap or cfg.vix_max_prior > 0):
        from RenTech.strategy_stack.cm_intraday_dip import CmIntradayConfig, build_regime_context

        regime = build_regime_context(
            master,
            spy_intra,
            CmIntradayConfig(
                require_spy_above_vwap=cfg.require_spy_above_vwap,
                vix_max_prior=cfg.vix_max_prior,
            ),
        )
    trades, port_r = simulate_atr_breakout(panels, cfg, regime=regime)
    port = pd.Series(port_r, index=master, dtype=np.float64)
    if return_start is not None:
        port = port.loc[port.index >= return_start]
    return port, pd.DataFrame([t.__dict__ for t in trades])
