"""
Intraday **same-session** MA slope top-N rotation on 5-minute (or N-minute) bars.

Rules
-----
* MA / slope periods are **bar counts** on intraday bars (e.g. EMA10 = 10×5m ≈ 50 min).
* Prior session 5m closes feed the MA (continuous series across days).
* **Trading window (default):** enter minutes **30–90** after 09:30 ET (5m bars 5–17).
* **Hold to MOC** (default): keep positions from entry through session close; no overnight.
* Optional ``hold_to_moc=False`` + ``session_max_hold_bar`` to flat at ~90 min instead.
* First bar of session never inherits prior-day position.
"""

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.ma_slope_engine import RankMetric, compute_slope_rank_score
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames

HoldMode = Literal["each_bar", "once_per_session"]
StopMode = Literal["none", "atr_trail", "pct_trail", "fixed_pct_entry"]

# 09:30 ET, 5m bars: bar 5 ≈ 10:00 entry exec, bar 17 ≈ 11:00 latest entry signal.
DEFAULT_SESSION_ENTRY_BAR_MIN = 5
DEFAULT_SESSION_ENTRY_BAR_MAX = 17
DEFAULT_SESSION_MAX_HOLD_BAR = 17  # used only when hold_to_moc=False
DEFAULT_ENTRY_BAR = 10


@dataclass(frozen=True)
class MaSlopeIntradayDayTradeConfig:
    ma_type: str = "ema"
    fast_period: int = 10
    slow_period: int = 50
    fast_lookback: int = 10
    slow_lookback: int = 5
    entry_slope_min: float = 0.0
    price_above_ma: bool = True
    rank_metric: RankMetric = "dual_product"
    min_score: float = 0.0
    hold_mode: HoldMode = "once_per_session"
    # Bar index within session for once_per_session entry (0 = 09:30 bar)
    entry_bar: int = DEFAULT_ENTRY_BAR
    # RTH 5m: entry only in minutes 30–90 after open (bars 5–17)
    session_entry_bar_min: int = DEFAULT_SESSION_ENTRY_BAR_MIN
    session_entry_bar_max: int = DEFAULT_SESSION_ENTRY_BAR_MAX
    hold_to_moc: bool = True
    session_max_hold_bar: int = DEFAULT_SESSION_MAX_HOLD_BAR  # exit early when hold_to_moc=False
    # Skip first N bars each session before entry (each_bar mode)
    skip_bars_per_session: int = 0
    rebalance_every_bars: int = 6
    # Per-name intraday stops (no re-leverage after stop; cash until session end)
    stop_mode: StopMode = "none"
    atr_period: int = 14
    atr_multiplier: float = 2.0
    pct_trail_stop: float = 0.02
    fixed_stop_pct: float = 0.015


@dataclass
class MaSlopeIntradayDayTrade:
    config: MaSlopeIntradayDayTradeConfig = MaSlopeIntradayDayTradeConfig()

    def build_score_panel(self, equity_dict: Dict[str, pd.DataFrame]) -> pd.DataFrame:
        cfg = self.config
        score_pan: list[pd.Series] = []
        for t, df in sorted(equity_dict.items()):
            if df is None or df.empty or "close" not in df.columns:
                continue
            idx = pd.to_datetime(df.index).tz_localize(None)
            close = df["close"].astype(np.float64)
            close.index = idx
            scored = compute_slope_rank_score(
                close.sort_index(),
                ma_type=cfg.ma_type,  # type: ignore[arg-type]
                fast_period=cfg.fast_period,
                slow_period=cfg.slow_period,
                fast_lookback=cfg.fast_lookback,
                slow_lookback=cfg.slow_lookback,
                slope_method="pct",
                entry_slope_min=cfg.entry_slope_min,
                price_above_ma=cfg.price_above_ma,
                rank_metric=cfg.rank_metric,
            )
            score_pan.append(scored["rank_score"].rename(t))
        if not score_pan:
            raise ValueError("no valid tickers")
        score_df, _ = _align_panel_frames(score_pan)
        return score_df.sort_index()

    @staticmethod
    def _session_key(index: pd.DatetimeIndex) -> pd.Series:
        idx = pd.to_datetime(index)
        return pd.Series(idx.normalize(), index=idx)

    @staticmethod
    def _last_bar_mask(index: pd.DatetimeIndex) -> pd.Series:
        ses = MaSlopeIntradayDayTrade._session_key(index)
        return ses != ses.shift(-1)

    @staticmethod
    def _session_bar_index(index: pd.DatetimeIndex) -> pd.Series:
        """0-based bar number within each session (resets at 09:30)."""
        idx = pd.to_datetime(index)
        ses = pd.Series(idx.normalize(), index=idx)
        out = pd.Series(0, index=idx, dtype=np.int64)
        prev_ses = None
        n = 0
        for i, dt in enumerate(idx):
            s = ses.iloc[i]
            if prev_ses is None or s != prev_ses:
                n = 0
                prev_ses = s
            else:
                n += 1
            out.iloc[i] = n
        return out

    @staticmethod
    def _exec_time_et(entry_bar: int) -> str:
        mins = 9 * 60 + 30 + (int(entry_bar) + 1) * 5
        h, m = divmod(mins, 60)
        return f"{h:02d}:{m:02d}"

    def _clamped_entry_bar(self) -> int:
        cfg = self.config
        lo = int(cfg.session_entry_bar_min)
        hi = int(cfg.session_entry_bar_max)
        return int(max(lo, min(int(cfg.entry_bar) + int(cfg.skip_bars_per_session), hi)))

    def _apply_session_window(self, w_df: pd.DataFrame) -> pd.DataFrame:
        """Entry-window hygiene; optional early flat when not holding to MOC."""
        cfg = self.config
        bar_in_ses = self._session_bar_index(w_df.index)
        lo = int(cfg.session_entry_bar_min)
        out = w_df.copy()
        if not cfg.hold_to_moc:
            out.loc[bar_in_ses > int(cfg.session_max_hold_bar)] = 0.0
        if cfg.hold_mode == "each_bar":
            out.loc[bar_in_ses < lo] = 0.0
            if not cfg.hold_to_moc:
                out.loc[bar_in_ses > int(cfg.session_max_hold_bar)] = 0.0
        return out

    @staticmethod
    def _first_bar_mask(index: pd.DatetimeIndex) -> pd.Series:
        ses = MaSlopeIntradayDayTrade._session_key(index)
        return ses != ses.shift(1)

    def target_weights(self, score_df: pd.DataFrame, top_n: int) -> pd.DataFrame:
        """Top-N equal weight from scores; flat after session_max_hold_bar and at close."""
        if top_n <= 0:
            raise ValueError("top_n must be > 0")
        cfg = self.config
        tickers = list(score_df.columns)
        score_arr = score_df.to_numpy(dtype=np.float64)
        n_bars, n_t = score_arr.shape
        w = np.zeros((n_bars, n_t), dtype=np.float64)
        ses = self._session_key(score_df.index)
        first = self._first_bar_mask(score_df.index)
        last = self._last_bar_mask(score_df.index)

        def _assign_topn(i: int) -> None:
            sc = score_arr[i]
            valid = np.isfinite(sc) & (sc > cfg.min_score)
            idx_valid = np.nonzero(valid)[0]
            if idx_valid.size == 0:
                return
            k = min(int(top_n), int(idx_valid.size))
            sub = sc[idx_valid]
            pick_local = np.argpartition(-sub, k - 1)[:k]
            chosen = idx_valid[pick_local]
            w[i, chosen] = 1.0 / float(k)

        if cfg.hold_mode == "each_bar":
            for i in range(n_bars):
                _assign_topn(i)
        elif cfg.hold_mode == "once_per_session":
            entry_bar = self._clamped_entry_bar()
            for i in range(n_bars):
                if not bool(first.iloc[i]):
                    continue
                j = i + entry_bar
                session_end = i
                while session_end + 1 < n_bars and ses.iloc[session_end + 1] == ses.iloc[i]:
                    session_end += 1
                j = min(j, session_end)
                _assign_topn(j)
                hold_end = session_end if cfg.hold_to_moc else min(session_end, i + int(cfg.session_max_hold_bar))
                row = w[j].copy()
                for k in range(j, hold_end):
                    w[k] = row
        elif cfg.hold_mode == "every_n_bars":
            step = max(1, int(cfg.rebalance_every_bars))
            lo = int(cfg.session_entry_bar_min)
            hi = int(cfg.session_entry_bar_max)
            bar_in_ses = self._session_bar_index(score_df.index)
            for i in range(n_bars):
                bis = int(bar_in_ses.iloc[i])
                if bis < lo or bis > hi:
                    continue
                if bis == lo or (bis - lo) % step == 0:
                    _assign_topn(i)
        else:
            raise ValueError(f"unknown hold_mode: {cfg.hold_mode!r}")

        w_df = pd.DataFrame(w, index=score_df.index, columns=tickers)
        w_df = self._apply_session_window(w_df)
        w_df.loc[last] = 0.0

        return w_df

    def _apply_intraday_stops(
        self,
        target_weights: pd.DataFrame,
        equity_dict: Dict[str, pd.DataFrame],
    ) -> pd.DataFrame:
        """Per-name stop between entry and session close; no same-day re-entry."""
        cfg = self.config
        if cfg.stop_mode == "none":
            return target_weights

        tickers = list(target_weights.columns)
        close_pan: list[pd.Series] = []
        high_pan: list[pd.Series] = []
        low_pan: list[pd.Series] = []
        for t in tickers:
            df = equity_dict.get(t)
            if df is None or df.empty:
                continue
            idx = pd.to_datetime(df.index).tz_localize(None)
            dfx = df.sort_index().copy()
            dfx.index = idx
            c = dfx["close"].astype(np.float64)
            h = dfx["high"].astype(np.float64) if "high" in dfx.columns else c
            lo = dfx["low"].astype(np.float64) if "low" in dfx.columns else c
            close_pan.append(c.rename(t))
            high_pan.append(h.rename(t))
            low_pan.append(lo.rename(t))

        close_df, _ = _align_panel_frames(close_pan)
        high_df, _ = _align_panel_frames(high_pan)
        low_df, _ = _align_panel_frames(low_pan)
        master = target_weights.index.sort_values()
        close_df = close_df.reindex(master).ffill()
        high_df = high_df.reindex(master).ffill()
        low_df = low_df.reindex(master).ffill()
        tw = target_weights.reindex(master).fillna(0.0).to_numpy(dtype=np.float64)
        c_arr = close_df.to_numpy(dtype=np.float64)
        h_arr = high_df.to_numpy(dtype=np.float64)
        lo_arr = low_df.to_numpy(dtype=np.float64)
        n_bars, n_t = tw.shape
        first = self._first_bar_mask(master).to_numpy(dtype=bool)

        atr_arr = None
        if cfg.stop_mode == "atr_trail":
            atr_arr = np.zeros((n_bars, n_t), dtype=np.float64)
            for j, t in enumerate(tickers):
                if t not in close_df.columns:
                    continue
                c = close_df[t]
                h = high_df[t]
                prev = c.shift(1)
                tr = pd.concat([h - c, (h - prev).abs(), (c - prev).abs()], axis=1).max(axis=1)
                atr_arr[:, j] = tr.rolling(cfg.atr_period, min_periods=cfg.atr_period).mean().to_numpy()

        w_out = tw.copy()
        entry_px = np.full(n_t, np.nan)
        peak_px = np.full(n_t, np.nan)
        stopped = np.zeros(n_t, dtype=bool)

        for i in range(n_bars):
            if i == 0 or first[i]:
                stopped[:] = False
                entry_px[:] = np.nan
                peak_px[:] = np.nan

            for j in range(n_t):
                tgt = tw[i, j]
                if tgt <= 0.0:
                    w_out[i, j] = 0.0
                    if stopped[j]:
                        stopped[j] = False
                    continue

                if stopped[j]:
                    w_out[i, j] = 0.0
                    continue

                new_pos = not np.isfinite(entry_px[j]) or (i > 0 and tw[i - 1, j] <= 0.0)
                if new_pos:
                    entry_px[j] = c_arr[i, j]
                    peak_px[j] = h_arr[i, j]
                    stopped[j] = False

                peak_px[j] = max(peak_px[j], h_arr[i, j])
                chk_px = lo_arr[i, j] if np.isfinite(lo_arr[i, j]) else c_arr[i, j]
                hit = False
                if cfg.stop_mode == "pct_trail":
                    hit = chk_px < peak_px[j] * (1.0 - float(cfg.pct_trail_stop))
                elif cfg.stop_mode == "atr_trail" and atr_arr is not None:
                    stop_px = peak_px[j] - float(cfg.atr_multiplier) * atr_arr[i, j]
                    hit = np.isfinite(stop_px) and chk_px < stop_px
                elif cfg.stop_mode == "fixed_pct_entry":
                    hit = chk_px < entry_px[j] * (1.0 - float(cfg.fixed_stop_pct))

                w_out[i, j] = tgt
                if hit:
                    stopped[j] = True

        return pd.DataFrame(w_out, index=master, columns=tickers)

    def generate_returns(
        self,
        equity_dict: Dict[str, pd.DataFrame],
        top_n: int = 10,
        *,
        return_start: pd.Timestamp | None = None,
        score_df: pd.DataFrame | None = None,
        verbose: bool = True,
    ) -> pd.Series:
        ret_pan: list[pd.Series] = []
        for t, df in sorted(equity_dict.items()):
            if df is None or df.empty or "ret" not in df.columns:
                continue
            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))
        if not ret_pan:
            raise ValueError("no return series")
        ret_df, _ = _align_panel_frames(ret_pan)
        master = ret_df.index.sort_values()
        ret_df = ret_df.reindex(master).fillna(0.0)

        if score_df is None:
            score_df = self.build_score_panel(equity_dict).reindex(columns=ret_df.columns)
        else:
            score_df = score_df.reindex(columns=ret_df.columns)
        score_df = score_df.reindex(master)

        target_w = self.target_weights(score_df, top_n)
        target_w = self._apply_intraday_stops(target_w, equity_dict)
        exec_w = target_w.shift(1).fillna(0.0)
        first = self._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 = (w * r).sum(axis=1)

        if verbose:
            w_sum = w.sum(axis=1)
            bar_invested = (w_sum > 0.01).mean()
            when_inv = w_sum > 0.01
            avg_k = float((w[when_inv] > 0).sum(axis=1).mean()) if when_inv.any() else 0.0
            eb = self._clamped_entry_bar()
            cfg = self.config
            hold_lbl = "MOC" if cfg.hold_to_moc else f"flat bar {cfg.session_max_hold_bar}"
            stop_lbl = cfg.stop_mode if cfg.stop_mode != "none" else "no stop"
            if cfg.stop_mode == "atr_trail":
                stop_lbl = f"atr{cfg.atr_multiplier:g}x"
            elif cfg.stop_mode == "pct_trail":
                stop_lbl = f"pct{int(cfg.pct_trail_stop * 100)}"
            elif cfg.stop_mode == "fixed_pct_entry":
                stop_lbl = f"fix{int(cfg.fixed_stop_pct * 1000) / 10:g}%"
            print(
                f"  Intraday day-trade top-{top_n}: "
                f"EMA{cfg.fast_period}/{cfg.slow_period} on 5m | "
                f"entry window bars {cfg.session_entry_bar_min}–{cfg.session_entry_bar_max} "
                f"(~{self._exec_time_et(cfg.session_entry_bar_min - 1)}–{self._exec_time_et(cfg.session_entry_bar_max)}) | "
                f"entry bar {eb} (~{self._exec_time_et(eb)}) | hold {hold_lbl} | stop {stop_lbl} | "
                f"invested {bar_invested * 100:.1f}% of bars | avg names ≈ {avg_k:.1f}",
                flush=True,
            )

        return pd.Series(port_r, index=master, name="intraday_daytrade_ret", dtype=np.float64).pipe(
            lambda s: s.loc[s.index >= return_start] if return_start is not None else s
        )
