"""
Enhanced intraday MA-slope day-trade: filters, exits, sizing, realism.

Built on ``ma_slope_intraday_daytrade`` defaults (enter 30–90 min, hold MOC).
"""

from __future__ import annotations

from dataclasses import dataclass, replace
from typing import Dict, Literal

import numpy as np
import pandas as pd

from RenTech.strategy_stack.ma_slope_engine import compute_slope_rank_score
from RenTech.strategy_stack.ma_slope_intraday_daytrade import (
    DEFAULT_ENTRY_BAR,
    DEFAULT_SESSION_ENTRY_BAR_MAX,
    DEFAULT_SESSION_ENTRY_BAR_MIN,
    MaSlopeIntradayDayTrade,
    MaSlopeIntradayDayTradeConfig,
)
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames

WeightMode = Literal["equal", "score_prop"]
SpyFilter = Literal["none", "dual_slope", "vwap"]
ExitMode = Literal["none", "slope_flip", "price_cross"]
HoldMode = Literal["once_per_session", "confirm_entry"]


@dataclass(frozen=True)
class EnhancedIntradayConfig(MaSlopeIntradayDayTradeConfig):
    weight_mode: WeightMode = "equal"
    max_weight_per_name: float = 0.25
    spy_filter: SpyFilter = "none"
    exit_mode: ExitMode = "none"
    hold_mode: HoldMode = "once_per_session"  # type: ignore[assignment]
    confirm_lag_bars: int = 4
    require_above_vwap: bool = False
    require_or_break: bool = False
    require_price_rising_confirm: bool = False  # confirm bar close > screen bar close
    min_rank_score_pct_gap: float = 0.0  # min (score_N - score_N+1) / score_N to take N names
    or_bars: int = 6
    session_stop_pct: float = 0.0
    spy_gross_scale: bool = False
    sector_cap: int = 0
    slippage_bps: float = 0.0
    min_price: float = 0.0
    min_adv_usd: float = 0.0


@dataclass
class EnhancedPanels:
    score: pd.DataFrame
    fast_slope: pd.DataFrame
    fast_ma: pd.DataFrame
    close: pd.DataFrame
    vwap: pd.DataFrame
    or_high: pd.DataFrame
    sector: dict[str, str]


def load_sp500_parquet_tickers(data_dir) -> list[str]:
    from RenTech.strategy_stack.alpaca_minute_loader import DEFAULT_ALPACA_RTH_DIR, list_parquet_symbols
    from RenTech.strategy_stack.run_johansen_triplet_sp500 import load_sp500_sectors

    have = set(list_parquet_symbols(data_dir))
    sec = load_sp500_sectors()
    return sorted(t for t in sec["ticker"].astype(str).str.upper() if t in have)


def liquidity_filter_symbols(daily_dict: dict[str, pd.DataFrame], *, min_price: float, min_adv: float) -> set[str]:
    ok: set[str] = set()
    for t, df in daily_dict.items():
        if df is None or df.empty:
            continue
        c = df["close"].astype(float)
        v = df["volume"].astype(float) if "volume" in df.columns else pd.Series(1.0, index=df.index)
        if float(c.iloc[-1]) < min_price:
            continue
        adv = float((c * v).tail(20).mean())
        if adv < min_adv:
            continue
        ok.add(t)
    return ok


def _session_vwap(close: pd.Series, volume: pd.Series, index: pd.DatetimeIndex) -> pd.Series:
    ses = pd.Series(index.normalize(), index=index)
    df = pd.DataFrame({"c": close.values, "v": volume.reindex(index).fillna(0.0).values, "ses": ses.values}, index=index)
    df["pv"] = df["c"] * df["v"]
    df["cpv"] = df.groupby("ses")["pv"].cumsum()
    df["cv"] = df.groupby("ses")["v"].cumsum().replace(0, np.nan)
    return df["cpv"] / df["cv"]


def _session_or_high(high: pd.Series, index: pd.DatetimeIndex, or_bars: int) -> pd.Series:
    bar_in_ses = EnhancedIntradayEngine._session_bar_index(index)
    ses = pd.Series(index.normalize(), index=index)
    df = pd.DataFrame({"h": high.values, "ses": ses.values, "bar": bar_in_ses.values}, index=index)
    df.loc[df["bar"] >= int(or_bars), "h"] = np.nan
    orh = df.groupby("ses")["h"].transform("max")
    return orh.reindex(index)


class EnhancedIntradayEngine(MaSlopeIntradayDayTrade):
    config: EnhancedIntradayConfig = EnhancedIntradayConfig()  # type: ignore[assignment]

    @staticmethod
    def _session_bar_index(index: pd.DatetimeIndex) -> pd.Series:
        return MaSlopeIntradayDayTrade._session_bar_index(index)

    def build_panels(self, equity_dict: Dict[str, pd.DataFrame], sectors: dict[str, str] | None = None) -> EnhancedPanels:
        cfg = self.config
        score_pan, fs_pan, fm_pan, cl_pan, vw_pan, or_pan = [], [], [], [], [], []
        for t, df in sorted(equity_dict.items()):
            if df is None or df.empty:
                continue
            idx = pd.to_datetime(df.index).tz_localize(None)
            close = df["close"].astype(np.float64)
            close.index = idx
            vol = df["volume"].astype(np.float64) if "volume" in df.columns else pd.Series(1.0, index=idx)
            high = df["high"].astype(np.float64) if "high" in df.columns else close
            high.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,
                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))
            fs_pan.append(scored["fast_slope"].rename(t))
            fm_pan.append(scored["fast_ma"].rename(t))
            cl_pan.append(close.rename(t))
            vw_pan.append(_session_vwap(close, vol, idx).rename(t))
            or_pan.append(_session_or_high(high, idx, cfg.or_bars).rename(t))
        score_df, _ = _align_panel_frames(score_pan)
        fs_df, _ = _align_panel_frames(fs_pan)
        fm_df, _ = _align_panel_frames(fm_pan)
        cl_df, _ = _align_panel_frames(cl_pan)
        vw_df, _ = _align_panel_frames(vw_pan)
        or_df, _ = _align_panel_frames(or_pan)
        sec = sectors or {}
        return EnhancedPanels(score_df.sort_index(), fs_df.sort_index(), fm_df.sort_index(), cl_df.sort_index(), vw_df.sort_index(), or_df.sort_index(), sec)

    def _pick_topn(
        self,
        sc: np.ndarray,
        *,
        top_n: int,
        tickers: list[str],
        sectors: dict[str, str],
    ) -> tuple[np.ndarray, float]:
        cfg = self.config
        valid = np.isfinite(sc) & (sc > cfg.min_score)
        idx_valid = np.nonzero(valid)[0]
        if idx_valid.size == 0:
            return np.zeros(len(sc)), 0.0
        order = idx_valid[np.argsort(-sc[idx_valid])]
        eff_n = self._topn_with_rank_gap(sc, order, top_n)
        if eff_n <= 0:
            return np.zeros(len(sc)), 0.0
        chosen: list[int] = []
        sec_counts: dict[str, int] = {}
        cap = int(cfg.sector_cap)
        for j in order:
            if len(chosen) >= eff_n:
                break
            if cap > 0:
                sec = sectors.get(tickers[j], "Unknown")
                if sec_counts.get(sec, 0) >= cap:
                    continue
                sec_counts[sec] = sec_counts.get(sec, 0) + 1
            chosen.append(int(j))
        if not chosen:
            return np.zeros(len(sc)), 0.0
        w = np.zeros(len(sc), dtype=np.float64)
        sub_sc = sc[chosen]
        if cfg.weight_mode == "score_prop":
            raw = np.maximum(sub_sc, 0.0)
            s = float(raw.sum())
            weights = raw / s if s > 0 else np.ones(len(chosen)) / len(chosen)
        else:
            weights = np.ones(len(chosen)) / len(chosen)
        mx = float(cfg.max_weight_per_name)
        if mx > 0 and mx < 1.0:
            weights = np.minimum(weights, mx)
            if cfg.weight_mode == "score_prop":
                s = float(weights.sum())
                weights = weights / s if s > 0 else weights
        for jj, ww in zip(chosen, weights):
            w[jj] = ww
        return w, float(weights.sum())

    def _topn_with_rank_gap(self, sc: np.ndarray, order: np.ndarray, top_n: int) -> int:
        gap = float(self.config.min_rank_score_pct_gap)
        if gap <= 0 or order.size <= top_n:
            return min(top_n, int(order.size))
        k = min(top_n, int(order.size))
        while k > 0:
            if k >= order.size:
                return k
            s_hi = float(sc[order[k - 1]])
            s_lo = float(sc[order[k]])
            if not np.isfinite(s_hi) or s_hi <= 0:
                k -= 1
                continue
            rel = (s_hi - s_lo) / s_hi
            if rel >= gap:
                return k
            k -= 1
        return 0

    def _bar_valid_mask(
        self, i: int, panels: EnhancedPanels, spy_row: dict | None
    ) -> np.ndarray:
        cfg = self.config
        n_t = len(panels.score.columns)
        ok = np.ones(n_t, dtype=bool)
        if cfg.require_above_vwap:
            c = panels.close.iloc[i].to_numpy()
            v = panels.vwap.iloc[i].to_numpy()
            ok &= np.isfinite(c) & np.isfinite(v) & (c > v)
        if cfg.require_or_break:
            c = panels.close.iloc[i].to_numpy()
            o = panels.or_high.iloc[i].to_numpy()
            ok &= np.isfinite(c) & np.isfinite(o) & (c > o)
        if cfg.min_price > 0:
            c = panels.close.iloc[i].to_numpy()
            ok &= np.isfinite(c) & (c >= float(cfg.min_price))
        if spy_row is not None and cfg.spy_filter != "none":
            if not spy_row.get("ok", True):
                ok[:] = False
        return ok

    def _spy_session_ok(self, panels: EnhancedPanels, i: int, spy_ticker: str = "SPY") -> tuple[bool, float]:
        cfg = self.config
        if cfg.spy_filter == "none" and not cfg.spy_gross_scale:
            return True, 1.0
        if spy_ticker not in panels.score.columns:
            return True, 1.0
        sc = float(panels.score.iloc[i][spy_ticker])
        c = float(panels.close.iloc[i][spy_ticker])
        v = float(panels.vwap.iloc[i][spy_ticker])
        ok = True
        if cfg.spy_filter == "dual_slope":
            ok = np.isfinite(sc) and sc > 0
        elif cfg.spy_filter == "vwap":
            ok = np.isfinite(c) and np.isfinite(v) and c > v
        scale = 1.0
        if cfg.spy_gross_scale and np.isfinite(sc) and sc > 0:
            scale = float(np.clip(sc / 1e-6, 0.25, 1.0)) if sc < 1e-4 else float(np.clip(sc * 500, 0.25, 1.0))
            scale = float(np.clip(scale, 0.25, 1.0))
        return bool(ok), scale

    def target_weights_enhanced(self, panels: EnhancedPanels, top_n: int) -> pd.DataFrame:
        cfg = self.config
        score_df = panels.score
        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_at(i: int, scale: float = 1.0) -> np.ndarray:
            sc = score_arr[i].copy()
            mask = self._bar_valid_mask(i, panels, None)
            sc[~mask] = np.nan
            row, gross = self._pick_topn(sc, top_n=top_n, tickers=tickers, sectors=panels.sector)
            return row * scale if gross > 0 else row

        if cfg.hold_mode == "confirm_entry":
            lag = int(cfg.confirm_lag_bars)
            entry_bar = self._clamped_entry_bar()
            for i in range(n_bars):
                if not bool(first.iloc[i]):
                    continue
                j1 = min(i + entry_bar, n_bars - 1)
                j2 = min(j1 + lag, n_bars - 1)
                session_end = i
                while session_end + 1 < n_bars and ses.iloc[session_end + 1] == ses.iloc[i]:
                    session_end += 1
                spy_ok, scale = self._spy_session_ok(panels, j2)
                if not spy_ok:
                    continue
                sc1 = score_arr[j1].copy()
                sc2 = score_arr[j2].copy()
                m1 = self._bar_valid_mask(j1, panels, None)
                m2 = self._bar_valid_mask(j2, panels, None)
                sc1[~m1] = np.nan
                sc2[~m2] = np.nan
                row1, _ = self._pick_topn(sc1, top_n=top_n, tickers=tickers, sectors=panels.sector)
                row2, _ = self._pick_topn(sc2, top_n=top_n, tickers=tickers, sectors=panels.sector)
                both = (row1 > 0) & (row2 > 0)
                if cfg.require_price_rising_confirm:
                    c1 = panels.close.iloc[j1].to_numpy()
                    c2 = panels.close.iloc[j2].to_numpy()
                    both &= np.isfinite(c1) & np.isfinite(c2) & (c2 > c1)
                if not both.any():
                    continue
                sc2_f = sc2.copy()
                sc2_f[~both] = np.nan
                row, _ = self._pick_topn(sc2_f, top_n=top_n, tickers=tickers, sectors=panels.sector)
                row *= scale
                hold_end = session_end if cfg.hold_to_moc else min(session_end, i + int(cfg.session_max_hold_bar))
                for k in range(j2, hold_end):
                    w[k] = row
        else:
            entry_bar = self._clamped_entry_bar()
            for i in range(n_bars):
                if not bool(first.iloc[i]):
                    continue
                j = min(i + entry_bar, n_bars - 1)
                session_end = i
                while session_end + 1 < n_bars and ses.iloc[session_end + 1] == ses.iloc[i]:
                    session_end += 1
                spy_ok, scale = self._spy_session_ok(panels, j)
                if not spy_ok:
                    continue
                row = assign_at(j, scale)
                hold_end = session_end if cfg.hold_to_moc else min(session_end, i + int(cfg.session_max_hold_bar))
                for k in range(j, hold_end):
                    w[k] = row

        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_momentum_exits(
        self, tw: pd.DataFrame, panels: EnhancedPanels
    ) -> pd.DataFrame:
        cfg = self.config
        if cfg.exit_mode == "none":
            return tw
        master = tw.index.sort_values()
        arr = tw.reindex(master).fillna(0.0).to_numpy(dtype=np.float64)
        fs = panels.fast_slope.reindex(master).to_numpy(dtype=np.float64)
        cl = panels.close.reindex(master).to_numpy(dtype=np.float64)
        fm = panels.fast_ma.reindex(master).to_numpy(dtype=np.float64)
        n_bars, n_t = arr.shape
        first = self._first_bar_mask(master).to_numpy(dtype=bool)
        stopped = np.zeros(n_t, dtype=bool)
        out = arr.copy()
        for i in range(n_bars):
            if i == 0 or first[i]:
                stopped[:] = False
            for j in range(n_t):
                tgt = arr[i, j]
                if tgt <= 0:
                    out[i, j] = 0.0
                    continue
                if stopped[j]:
                    out[i, j] = 0.0
                    continue
                hit = False
                if cfg.exit_mode == "slope_flip":
                    hit = np.isfinite(fs[i, j]) and fs[i, j] < 0
                elif cfg.exit_mode == "price_cross":
                    hit = np.isfinite(cl[i, j]) and np.isfinite(fm[i, j]) and cl[i, j] < fm[i, j]
                out[i, j] = tgt
                if hit:
                    stopped[j] = True
        return pd.DataFrame(out, index=master, columns=tw.columns)

    def _apply_session_stop(self, port_r: pd.Series, exec_w: pd.DataFrame) -> pd.Series:
        pct = float(self.config.session_stop_pct)
        if pct <= 0:
            return port_r
        ses = self._session_key(port_r.index)
        first = self._first_bar_mask(port_r.index)
        out = port_r.to_numpy(dtype=np.float64).copy()
        blocked = False
        ses_eq = 1.0
        for i in range(len(out)):
            if bool(first.iloc[i]):
                blocked = False
                ses_eq = 1.0
            if blocked:
                out[i] = 0.0
                continue
            out[i] = port_r.iloc[i]
            ses_eq *= 1.0 + out[i]
            if ses_eq / 1.0 - 1.0 <= -pct:
                blocked = True
        return pd.Series(out, index=port_r.index)

    def _apply_slippage(self, port_r: pd.Series, exec_w: pd.DataFrame) -> pd.Series:
        bps = float(self.config.slippage_bps)
        if bps <= 0:
            return port_r
        gross = exec_w.sum(axis=1).astype(float)
        dg = gross.diff().fillna(gross).clip(lower=0.0)
        cost = dg * (bps / 10_000.0)
        return port_r - cost

    def run(
        self,
        equity_dict: Dict[str, pd.DataFrame],
        top_n: int,
        *,
        panels: EnhancedPanels | None = None,
        return_start: pd.Timestamp | None = None,
    ) -> pd.Series:
        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)
        if panels is None:
            panels = self.build_panels(equity_dict)
        panels = EnhancedPanels(
            panels.score.reindex(master).reindex(columns=ret_df.columns),
            panels.fast_slope.reindex(master).reindex(columns=ret_df.columns),
            panels.fast_ma.reindex(master).reindex(columns=ret_df.columns),
            panels.close.reindex(master).reindex(columns=ret_df.columns),
            panels.vwap.reindex(master).reindex(columns=ret_df.columns),
            panels.or_high.reindex(master).reindex(columns=ret_df.columns),
            panels.sector,
        )
        target_w = self.target_weights_enhanced(panels, top_n)
        target_w = self._apply_momentum_exits(target_w, panels)
        exec_w = target_w.shift(1).fillna(0.0)
        exec_w.loc[self._first_bar_mask(master)] = 0.0
        port_r = pd.Series((exec_w.to_numpy() * ret_df.to_numpy()).sum(axis=1), index=master)
        port_r = self._apply_session_stop(port_r, exec_w)
        port_r = self._apply_slippage(port_r, exec_w)
        if return_start is not None:
            port_r = port_r.loc[port_r.index >= return_start]
        return port_r


def baseline_enhanced_config() -> EnhancedIntradayConfig:
    return EnhancedIntradayConfig()


def metrics_daily(r: pd.Series) -> dict:
    from RenTech.strategy_stack.alpaca_minute_loader import compound_intraday_to_daily

    ds = compound_intraday_to_daily(r).dropna()
    if len(ds) < 2:
        return {}
    eq = (1 + ds).cumprod()
    years = len(ds) / 252.0
    sd = float(ds.std(ddof=1))
    return {
        "total_return_pct": float((eq.iloc[-1] - 1) * 100),
        "cagr_pct": float((eq.iloc[-1] ** (1 / years) - 1) * 100) if years > 0 else float("nan"),
        "max_dd_pct": float((eq / eq.cummax() - 1).min() * 100),
        "sharpe": float(ds.mean() / sd * np.sqrt(252)) if sd > 1e-12 else float("nan"),
        "n_sessions": int(len(ds)),
    }
