"""
Cross-sectional **MA slope momentum** — rank universe, hold top-N equal weight.

Monthly (or weekly) rebalance: score = dual EMA slope strength (stage-4 SPY winner
logic), pick the ``top_n`` names with highest ``rank_score``.
"""

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 (
    MaType,
    RankMetric,
    SlopeMethod,
    compute_slope_rank_score,
)
from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames

RebalanceFreq = Literal["monthly", "weekly"]
StopMode = Literal["none", "atr_trail", "pct_trail", "portfolio_dd"]


@dataclass(frozen=True)
class MaSlopeCrossSectionalConfig:
    ma_type: MaType = "ema"
    fast_period: int = 10
    slow_period: int = 50
    fast_lookback: int = 10
    slow_lookback: int = 5
    slope_method: SlopeMethod = "pct"
    entry_slope_min: float = 0.0
    price_above_ma: bool = True
    rank_metric: RankMetric = "dual_product"
    rebalance: RebalanceFreq = "monthly"
    cash_annual_yield: float = 0.04
    bars_per_year: int = 252
    # Per-name ATR chandelier between rebalances (canonical: 2× ATR trail)
    stop_mode: StopMode = "atr_trail"
    atr_period: int = 14
    atr_multiplier: float = 2.0
    pct_trail_stop: float = 0.12
    portfolio_dd_stop: float = 0.10


@dataclass
class MaSlopeCrossSectional:
    """Long-only top-N rotation by MA slope rank score."""

    config: MaSlopeCrossSectionalConfig = MaSlopeCrossSectionalConfig()

    def _rebalance_index(self, score_df: pd.DataFrame) -> pd.DatetimeIndex:
        idx = pd.to_datetime(score_df.index).sort_values()
        if self.config.rebalance == "monthly":
            bars_per_day = pd.Series(1, index=idx).groupby(idx.normalize()).sum()
            if float(bars_per_day.median()) > 1.5:
                # Intraday panels: last session bar each month (BME → midnight breaks reindex)
                last = score_df.loc[idx].groupby([idx.year, idx.month]).tail(1)
                return pd.DatetimeIndex(last.index)
            return score_df.resample("BME").last().index
        if self.config.rebalance == "weekly":
            bars_per_day = pd.Series(1, index=idx).groupby(idx.normalize()).sum()
            if float(bars_per_day.median()) > 1.5:
                last = score_df.loc[idx].groupby(idx.to_period("W-FRI")).tail(1)
                return pd.DatetimeIndex(last.index)
            return score_df.resample("W-FRI").last().index
        raise ValueError(f"unknown rebalance: {self.config.rebalance!r}")

    def build_score_panel(self, equity_dict: Dict[str, pd.DataFrame]) -> pd.DataFrame:
        """Wide panel of ``rank_score`` per ticker (DatetimeIndex × tickers)."""
        return self.build_slope_panels(equity_dict)[0]

    def build_slope_panels(
        self, equity_dict: Dict[str, pd.DataFrame]
    ) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:
        """Return (rank_score, fast_slope, slow_slope) wide panels."""
        if not equity_dict:
            raise ValueError("equity_dict cannot be empty")
        cfg = self.config
        score_pan: list[pd.Series] = []
        fast_pan: list[pd.Series] = []
        slow_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,
                fast_period=cfg.fast_period,
                slow_period=cfg.slow_period,
                fast_lookback=cfg.fast_lookback,
                slow_lookback=cfg.slow_lookback,
                slope_method=cfg.slope_method,
                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))
            fast_pan.append(scored["fast_slope"].rename(t))
            slow_pan.append(scored["slow_slope"].rename(t))
        if not score_pan:
            raise ValueError("no valid tickers with close prices")
        score_df, _ = _align_panel_frames(score_pan)
        fast_df, _ = _align_panel_frames(fast_pan)
        slow_df, _ = _align_panel_frames(slow_pan)
        return score_df.sort_index(), fast_df.sort_index(), slow_df.sort_index()

    def _topn_weights(
        self,
        score_df: pd.DataFrame,
        ret_df: pd.DataFrame,
        top_n: int,
    ) -> tuple[pd.DataFrame, pd.DatetimeIndex, np.ndarray]:
        bm_index = self._rebalance_index(score_df)
        score_m = score_df.reindex(bm_index, method="ffill")
        score_arr = score_m.to_numpy(dtype=np.float64)
        M, N = score_arr.shape
        weights_m = np.zeros((M, N), dtype=np.float64)

        for m in range(M):
            sc = score_arr[m]
            valid = np.isfinite(sc) & (sc > 0.0)
            idx_valid = np.nonzero(valid)[0]
            if idx_valid.size == 0:
                continue
            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]
            order = np.argsort(-sub[pick_local])
            chosen = chosen[order]
            weights_m[m, chosen] = 1.0 / float(k)

        weights_m_df = pd.DataFrame(weights_m, index=bm_index, columns=score_df.columns)
        master_index = ret_df.index.sort_values()
        weights_d = weights_m_df.reindex(master_index).ffill().shift(1).fillna(0.0)
        return weights_d, bm_index, weights_m

    def _apply_stops(
        self,
        target_weights: pd.DataFrame,
        equity_dict: Dict[str, pd.DataFrame],
        bm_index: pd.DatetimeIndex,
    ) -> pd.DataFrame:
        """Zero stopped-out names between rebalances; cash until next rebalance."""
        cfg = self.config
        mode = cfg.stop_mode
        if mode == "none":
            return target_weights

        tickers = list(target_weights.columns)
        close_pan: list[pd.Series] = []
        high_pan: list[pd.Series] = []
        ret_pan: list[pd.Series] = []
        for t in tickers:
            df = equity_dict[t]
            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
            close_pan.append(c.rename(t))
            high_pan.append(h.rename(t))
            ret_pan.append(dfx["ret"].astype(np.float64).rename(t))

        close_df, _ = _align_panel_frames(close_pan)
        high_df, _ = _align_panel_frames(high_pan)
        ret_df, _ = _align_panel_frames(ret_pan)
        master = target_weights.index.sort_values()
        close_df = close_df.reindex(master).ffill()
        high_df = high_df.reindex(master).ffill()
        ret_df = ret_df.reindex(master).fillna(0.0)
        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)
        r_arr = ret_df.to_numpy(dtype=np.float64)
        n_d, n_t = tw.shape
        reb_set = {pd.Timestamp(x).normalize() for x in bm_index}
        daily_rf = float(cfg.cash_annual_yield) / float(cfg.bars_per_year)

        atr_arr = None
        if mode == "atr_trail":
            atr_arr = np.zeros((n_d, n_t), dtype=np.float64)
            for j, t in enumerate(tickers):
                c = close_df[t]
                h = high_df[t]
                prev = c.shift(1)
                tr_s = pd.concat([h - c, (h - prev).abs(), (c - prev).abs()], axis=1).max(axis=1)
                atr_arr[:, j] = tr_s.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)
        port_blocked = False
        port_block_next = False
        eq = 1.0
        peak_eq = 1.0

        for i in range(n_d):
            dt = pd.Timestamp(master[i]).normalize()
            if dt in reb_set:
                port_blocked = False
                port_block_next = False
                stopped[:] = False
                entry_px[:] = np.nan
                peak_px[:] = np.nan

            if mode == "portfolio_dd":
                port_blocked = port_block_next
                port_block_next = False

            if mode == "portfolio_dd" and port_blocked:
                w_out[i] = 0.0
                continue

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

                if dt in reb_set or not np.isfinite(entry_px[j]):
                    entry_px[j] = c_arr[i, j]
                    peak_px[j] = h_arr[i, j]
                    stopped[j] = False

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

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

                w_out[i, j] = tgt
                if hit:
                    stopped[j] = True  # flat from next session (still held through trigger day)

            # Do not renormalize to 100% — stopped names sit in cash (true de-risk).
            if mode == "portfolio_dd":
                day_r = (w_out[i] * r_arr[i]).sum() + max(0.0, 1.0 - w_out[i].sum()) * daily_rf
                eq *= 1.0 + day_r
                peak_eq = max(peak_eq, eq)
                if eq / peak_eq - 1.0 <= -cfg.portfolio_dd_stop:
                    port_block_next = True

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

    def generate_returns(
        self,
        equity_dict: Dict[str, pd.DataFrame],
        top_n: int = 10,
        *,
        verbose: bool = True,
    ) -> pd.Series:
        if top_n <= 0:
            raise ValueError("top_n must be > 0")

        tickers = sorted(equity_dict.keys())
        ret_pan: list[pd.Series] = []
        for t in tickers:
            df = equity_dict[t]
            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 in equity_dict")

        ret_df, _ = _align_panel_frames(ret_pan)
        score_df, _, _ = self.build_slope_panels(equity_dict)
        score_df = score_df.reindex(columns=ret_df.columns)
        master_index = ret_df.index.sort_values()
        ret_df = ret_df.reindex(master_index)
        ret_filled = ret_df.fillna(0.0).astype(np.float64)

        weights_d, bm_index, weights_m = self._topn_weights(score_df, ret_df, top_n)
        weights_d = self._apply_stops(weights_d, equity_dict, bm_index)
        weights_d = weights_d.reindex(master_index).fillna(0.0)

        daily_rf = float(self.config.cash_annual_yield) / float(self.config.bars_per_year)
        w = weights_d.to_numpy(dtype=np.float64)
        r = ret_filled.to_numpy(dtype=np.float64)
        core = (w * r).sum(axis=1)
        wsum = w.sum(axis=1)
        out = core + np.maximum(0.0, 1.0 - wsum) * daily_rf

        if verbose:
            M = len(bm_index)
            avg_k = float(np.mean([np.sum(weights_m[i] > 0) for i in range(M)])) if M else 0.0
            elig = score_df.reindex(bm_index, method="ffill").notna().sum(axis=1)
            avg_elig = float(elig.mean()) if len(elig) else 0.0
            print(
                f"  MA slope top-{top_n}: metric={self.config.rank_metric} | "
                f"rebal={self.config.rebalance} | avg held ≈ {avg_k:.1f} | "
                f"avg eligible ≈ {avg_elig:.0f} names/rebal",
                flush=True,
            )

        return pd.Series(out, index=master_index, name="ma_slope_topn_ret", dtype=np.float64)

    def generate_rebalance_log(
        self,
        equity_dict: Dict[str, pd.DataFrame],
        top_n: int = 10,
    ) -> pd.DataFrame:
        if top_n <= 0:
            raise ValueError("top_n must be > 0")

        score_df, fast_df, slow_df = self.build_slope_panels(equity_dict)
        tickers = list(score_df.columns)
        bm_index = self._rebalance_index(score_df)
        score_m = score_df.reindex(bm_index, method="ffill")
        fast_m = fast_df.reindex(bm_index, method="ffill")
        slow_m = slow_df.reindex(bm_index, method="ffill")
        score_arr = score_m.to_numpy(dtype=np.float64)
        fast_arr = fast_m.to_numpy(dtype=np.float64)
        slow_arr = slow_m.to_numpy(dtype=np.float64)
        master_index = score_df.index.sort_values()
        M, N = score_arr.shape
        if M < 1:
            return pd.DataFrame()

        rows: list[dict] = []
        prev_held: set[str] = set()
        for m in range(M):
            signal_dt = pd.Timestamp(bm_index[m]).normalize()
            sc = score_arr[m]
            valid = np.isfinite(sc) & (sc > 0.0)
            idx_valid = np.nonzero(valid)[0]
            if idx_valid.size == 0:
                prev_held = set()
                continue
            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]
            order = np.argsort(-sub[pick_local])
            chosen = chosen[order]

            held = {tickers[int(j)] for j in chosen}
            entries = held - prev_held
            exits = prev_held - held

            after = master_index[master_index > signal_dt]
            effective_dt = pd.Timestamp(after[0]).normalize() if len(after) else signal_dt
            period_end = (
                pd.Timestamp(bm_index[m + 1]).normalize() if m + 1 < M else pd.Timestamp(master_index[-1]).normalize()
            )

            for rank_i, j in enumerate(chosen, start=1):
                tkr = tickers[int(j)]
                rows.append(
                    {
                        "signal_date": signal_dt.strftime("%Y-%m-%d"),
                        "effective_date": effective_dt.strftime("%Y-%m-%d"),
                        "period_end_date": period_end.strftime("%Y-%m-%d"),
                        "ticker": tkr,
                        "weight": 1.0 / float(k),
                        "rank_score": float(sc[j]),
                        "fast_slope": float(fast_arr[m, j]),
                        "slow_slope": float(slow_arr[m, j]),
                        "rank": int(rank_i),
                        "is_new_entry": tkr in entries,
                        "is_exit_from_prior": False,
                        "top_n": int(k),
                    }
                )
            for tkr in sorted(exits):
                rows.append(
                    {
                        "signal_date": signal_dt.strftime("%Y-%m-%d"),
                        "effective_date": effective_dt.strftime("%Y-%m-%d"),
                        "period_end_date": period_end.strftime("%Y-%m-%d"),
                        "ticker": tkr,
                        "weight": 0.0,
                        "rank_score": float("nan"),
                        "fast_slope": float("nan"),
                        "slow_slope": float("nan"),
                        "rank": 0,
                        "is_new_entry": False,
                        "is_exit_from_prior": True,
                        "top_n": int(k),
                    }
                )
            prev_held = held

        return pd.DataFrame(rows)
