"""
S&P 500 **Momentum Index** sleeve — SPMO-style long-only rotation.

Core index rules plus optional **enhanced selection** filters (dual momentum, SMA trend,
rank gap, liquidity floors, sector-neutral picks, residual momentum, multi-horizon blend).
"""

from __future__ import annotations

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

import numpy as np
import pandas as pd

from RenTech.strategy_stack.multi_strategy_manager import _align_panel_frames, _distribute_slots

CapProxy = Literal["price", "mkt_cap", "equal"]
RebalanceFreq = Literal["semiannual", "monthly"]


def compute_momentum_features(
    close: pd.Series,
    volume: pd.Series | None = None,
    *,
    sma_window: int = 200,
) -> pd.DataFrame:
    """Per-ticker momentum, trend, and liquidity features (point-in-time)."""
    c = close.astype(np.float64).sort_index()
    ret = c.pct_change()
    raw_12 = c.shift(21) / c.shift(252) - 1.0
    raw_6 = c.shift(21) / c.shift(126) - 1.0
    raw_3 = c.shift(1) / c.shift(63) - 1.0
    vol_12 = ret.rolling(window=231, min_periods=126).std(ddof=1).shift(21)
    vol_6 = ret.rolling(window=105, min_periods=63).std(ddof=1).shift(21)
    risk_12 = raw_12 / vol_12.replace(0.0, np.nan)
    risk_6 = raw_6 / vol_6.replace(0.0, np.nan)
    td = ret.notna().rolling(231, min_periods=1).sum().shift(21)
    sma = c.rolling(int(sma_window), min_periods=int(sma_window)).mean() if sma_window > 0 else pd.Series(np.nan, index=c.index)
    if sma_window > 0:
        above_sma = (c > sma).astype(np.float64)
    else:
        above_sma = pd.Series(1.0, index=c.index, dtype=np.float64)
    roll_hi = c.rolling(252, min_periods=60).max()
    near_high = c / roll_hi.replace(0.0, np.nan)
    adv = pd.Series(np.nan, index=c.index, dtype=np.float64)
    if volume is not None and not volume.empty:
        v = volume.astype(np.float64).reindex(c.index).fillna(0.0)
        adv = (c * v).rolling(20, min_periods=10).mean()
    return pd.DataFrame(
        {
            "raw_12": raw_12,
            "raw_6": raw_6,
            "raw_3": raw_3,
            "risk_12": risk_12,
            "risk_6": risk_6,
            "vol_12": vol_12,
            "ret": ret,
            "trading_days": td,
            "above_sma": above_sma.astype(np.float64),
            "near_high": near_high.astype(np.float64),
            "adv_usd": adv,
        },
        index=c.index,
    )


def compute_momentum_score_panel(close: pd.Series) -> pd.DataFrame:
    """Legacy wrapper — 12-minus-1 risk-adjusted momentum."""
    f = compute_momentum_features(close, sma_window=0)
    return pd.DataFrame(
        {
            "raw_mom": f["raw_12"],
            "vol_mom_window": f["vol_12"],
            "risk_adj_mom": f["risk_12"],
            "ret": f["ret"],
        },
        index=close.index,
    )


def cross_sectional_zscore(values: np.ndarray) -> np.ndarray:
    out = np.full(values.shape, np.nan, dtype=np.float64)
    valid = np.isfinite(values)
    if valid.sum() < 2:
        return out
    mu = float(values[valid].mean())
    sigma = float(values[valid].std(ddof=0))
    if sigma <= 1e-14:
        return out
    z = (values - mu) / sigma
    out = np.clip(z, -3.0, 3.0)
    out[~valid] = np.nan
    return out


def cross_sectional_sp_momentum_scores(risk_adj_row: np.ndarray) -> np.ndarray:
    """S&P Appendix B: z-score → winsorize ±3 → momentum score transform."""
    z = cross_sectional_zscore(risk_adj_row)
    out = np.full(risk_adj_row.shape, np.nan, dtype=np.float64)
    valid = np.isfinite(z)
    scores = np.where(z > 0.0, 1.0 + z, np.where(z < 0.0, 1.0 / (1.0 - z), 1.0))
    out[valid] = scores[valid]
    return out


def blend_horizon_risk_adj(
    risk_6: np.ndarray,
    risk_12: np.ndarray,
    *,
    weight_6m: float,
    weight_12m: float,
) -> np.ndarray:
    w6, w12 = float(weight_6m), float(weight_12m)
    tot = w6 + w12
    if tot <= 1e-12:
        return risk_12.astype(np.float64)
    out = np.full(risk_12.shape, np.nan, dtype=np.float64)
    if w6 > 0:
        out = np.where(np.isfinite(risk_6), w6 * risk_6, out)
    if w12 > 0:
        out = np.where(np.isfinite(risk_12), np.where(np.isfinite(out), out + w12 * risk_12, w12 * risk_12), out)
    out = out / tot
    return out


def residualize_risk_row(
    raw_12: np.ndarray,
    vol_12: np.ndarray,
    beta: np.ndarray,
    spy_raw_12: float,
) -> np.ndarray:
    """Residual 12−1 momentum: stock raw return minus β × SPY raw return, vol-scaled."""
    if not np.isfinite(spy_raw_12):
        vol = np.where(vol_12 > 0, vol_12, np.nan)
        return raw_12 / vol
    res_raw = raw_12 - beta * spy_raw_12
    vol = np.where(vol_12 > 0, vol_12, np.nan)
    return res_raw / vol


def adjust_target_for_rank_gap(
    scores: np.ndarray,
    eligible: np.ndarray,
    target: int,
    min_gap: float,
) -> int:
    """Reduce ``target`` when top-N vs N+1 score gap is too narrow."""
    if min_gap <= 0.0 or target <= 0:
        return target
    idx = np.nonzero(eligible & np.isfinite(scores))[0]
    if idx.size == 0:
        return 0
    ordered = idx[np.argsort(-scores[idx])]
    max_k = min(int(target), int(ordered.size))
    for k in range(max_k, 0, -1):
        top = float(scores[ordered[k - 1]])
        if k >= ordered.size:
            return k
        nxt = float(scores[ordered[k]])
        if abs(top) <= 1e-12:
            return k
        gap = (top - nxt) / abs(top)
        if gap >= float(min_gap):
            return k
    return 1


def select_with_sp_buffer(
    scores: np.ndarray,
    eligible: np.ndarray,
    prev_held_idx: set[int],
    target: int,
    *,
    buffer_mult: float = 1.2,
) -> np.ndarray:
    n = scores.shape[0]
    mask = np.zeros(n, dtype=bool)
    if target <= 0:
        return mask
    idx_elig = np.nonzero(eligible & np.isfinite(scores))[0]
    if idx_elig.size == 0:
        return mask
    sub_scores = scores[idx_elig]
    order = np.argsort(-sub_scores)
    ranked_idx = idx_elig[order]
    auto_n = max(1, int(np.floor(0.8 * target)))
    buffer_n = max(auto_n, int(np.ceil(float(buffer_mult) * target)))
    band_lo, band_hi = auto_n, target
    selected: list[int] = []
    for i in range(min(auto_n, len(ranked_idx))):
        selected.append(int(ranked_idx[i]))
    sel_set = set(selected)
    for idx in ranked_idx[:buffer_n]:
        idx = int(idx)
        if idx in prev_held_idx and idx not in sel_set:
            selected.append(idx)
            sel_set.add(idx)
        if len(selected) >= target:
            break
    if len(selected) < target:
        for idx in ranked_idx[band_lo:band_hi]:
            idx = int(idx)
            if idx not in sel_set:
                selected.append(idx)
                sel_set.add(idx)
            if len(selected) >= target:
                break
    if len(selected) < target:
        for idx in ranked_idx:
            idx = int(idx)
            if idx not in sel_set:
                selected.append(idx)
                sel_set.add(idx)
            if len(selected) >= target:
                break
    for idx in selected[:target]:
        mask[idx] = True
    return mask


def select_sector_neutral_long(
    scores: np.ndarray,
    eligible: np.ndarray,
    sector_ids: np.ndarray,
    target: int,
) -> np.ndarray:
    """Within-sector momentum rank; ``target`` seats split across sectors."""
    n = scores.shape[0]
    mask = np.zeros(n, dtype=bool)
    if target <= 0:
        return mask
    sectors = sorted({str(sector_ids[j]) for j in range(n) if eligible[j]})
    if not sectors:
        return mask
    slots = _distribute_slots(int(target), len(sectors))
    for si, sec in enumerate(sectors):
        ks = int(slots[si])
        if ks <= 0:
            continue
        idxs = np.where(eligible & (sector_ids == sec))[0]
        if idxs.size == 0:
            continue
        take = min(ks, int(idxs.size))
        sub = scores[idxs].astype(np.float64)
        part = np.argpartition(-sub, take - 1)[:take]
        mask[idxs[part]] = True
    return mask


@dataclass(frozen=True)
class Sp500MomentumConfig:
    top_n: int = 0
    lookback: int = 252
    skip: int = 21
    rebalance: RebalanceFreq = "semiannual"
    semiannual_months: tuple[int, ...] = (2, 8)
    cap_proxy: CapProxy = "mkt_cap"
    max_single_weight: float = 0.09
    mkt_cap_multiplier: float = 3.0
    max_sector_weight: float = 0.0
    cash_annual_yield: float = 0.04
    min_mom_score: float = 0.0
    min_trading_days: int = 150
    use_sp_index_rules: bool = True
    use_buffer_rule: bool = True
    quintile_selection: bool = True
    # Enhanced selection (1–7)
    require_positive_raw_mom: bool = False
    min_raw_mom: float = 0.0
    above_sma_window: int = 0
    min_rank_score_pct_gap: float = 0.0
    min_market_cap_usd: float = 0.0
    min_adv_usd: float = 0.0
    sector_neutral_select: bool = False
    use_residual_momentum: bool = False
    residual_beta_window: int = 60
    mom_weight_6m: float = 0.0
    mom_weight_12m: float = 1.0
    score_weight_power: float = 1.0
    selection_filters: list[str] = field(default_factory=list)
    # Ride-the-rockets knobs
    require_positive_raw_6: bool = False
    require_positive_raw_3: bool = False
    use_raw_momentum_rank: bool = False
    near_high_min_frac: float = 0.0
    hold_buffer_mult: float = 0.0
    weight_by_score_only: bool = False


def enhanced_selection_defaults() -> dict:
    """Preset enabling filters 1–7 with sensible defaults."""
    return {
        "require_positive_raw_mom": True,
        "min_raw_mom": 0.0,
        "above_sma_window": 200,
        "min_rank_score_pct_gap": 0.05,
        "min_market_cap_usd": 1_000_000_000.0,
        "min_adv_usd": 5_000_000.0,
        "sector_neutral_select": True,
        "max_sector_weight": 0.25,
        "use_residual_momentum": True,
        "mom_weight_6m": 0.3,
        "mom_weight_12m": 0.7,
        "selection_filters": [
            "dual_momentum",
            "above_sma200",
            "rank_gap",
            "liquidity",
            "sector_neutral",
            "residual_mom",
            "multi_horizon",
        ],
    }


def best_of_best_defaults(*, top_n: int = 25) -> dict:
    """
    Concentrated **best-of-best** book: top ``top_n`` by S&P momentum score after
    dual-momentum + SMA200 filters (no sector-neutral quota; lighter liquidity).
    """
    return {
        "top_n": int(top_n),
        "quintile_selection": False,
        "use_buffer_rule": False,
        "require_positive_raw_mom": True,
        "above_sma_window": 200,
        "use_residual_momentum": True,
        "mom_weight_6m": 0.3,
        "mom_weight_12m": 0.7,
        "min_market_cap_usd": 500_000_000.0,
        "min_adv_usd": 2_000_000.0,
        "min_rank_score_pct_gap": 0.0,
        "sector_neutral_select": False,
        "max_sector_weight": 0.35,
        "score_weight_power": 1.25,
        "selection_filters": ["best_of_best", f"top_{int(top_n)}"],
    }


def ride_rockets_defaults(*, top_n: int = 25) -> dict:
    """Concentrated monthly book: ride relative winners until thrust fades."""
    return {
        "top_n": int(top_n),
        "quintile_selection": False,
        "use_buffer_rule": False,
        "rebalance": "monthly",
        "require_positive_raw_mom": True,
        "above_sma_window": 150,
        "use_residual_momentum": False,
        "mom_weight_6m": 0.3,
        "mom_weight_12m": 0.7,
        "min_market_cap_usd": 0.0,
        "min_adv_usd": 5_000_000.0,
        "min_rank_score_pct_gap": 0.0,
        "sector_neutral_select": False,
        "max_sector_weight": 0.0,
        "max_single_weight": 0.15,
        "score_weight_power": 1.25,
        "use_sp_index_rules": True,
        "cap_proxy": "mkt_cap",
        "require_positive_raw_6": False,
        "require_positive_raw_3": False,
        "use_raw_momentum_rank": False,
        "near_high_min_frac": 0.0,
        "hold_buffer_mult": 0.0,
        "weight_by_score_only": False,
        "selection_filters": ["ride_rockets", f"top_{int(top_n)}"],
    }


def ride_rockets_champ_defaults(*, top_n: int = 15) -> dict:
    """
    Empirically tuned champ after the combined grid:

    - Top 15 + near-52w-high + 6m-fade kill won Sharpe / DD in the champ grid.
    - Tighter books (top 10–12) with *both* gates over-filtered vs prior singles.
    - For return-max, prefer kill-only (see ``ride_rockets_champ_catalog`` ablations).
    """
    n = int(top_n)
    cap = 0.20 if n <= 10 else (0.18 if n <= 12 else 0.15)
    return {
        **ride_rockets_defaults(top_n=n),
        "near_high_min_frac": 0.95,
        "require_positive_raw_6": True,
        "max_single_weight": cap,
        "score_weight_power": 1.35 if n <= 12 else 1.25,
        "selection_filters": ["ride_rockets", "champ", f"top_{n}"],
    }


def ride_rockets_champ_catalog() -> list[tuple[str, str, dict]]:
    """Small grid around the champ recipe (top 10 / 12 / 15 ± ablations)."""
    return [
        (
            "champ_top15",
            "Default champ: top 15 + near-52w-high + 6m-fade kill",
            ride_rockets_champ_defaults(top_n=15),
        ),
        (
            "champ_top10",
            "Tighter: top 10 + near-52w-high + 6m-fade kill",
            ride_rockets_champ_defaults(top_n=10),
        ),
        (
            "champ_top12",
            "Mid: top 12 + near-52w-high + 6m-fade kill",
            ride_rockets_champ_defaults(top_n=12),
        ),
        (
            "champ_top12_near_only",
            "Ablation: top 12 + near-high, no 6m kill",
            {
                **ride_rockets_champ_defaults(top_n=12),
                "require_positive_raw_6": False,
                "selection_filters": ["ride_rockets", "champ", "near_only", "top_12"],
            },
        ),
        (
            "champ_top12_kill_only",
            "Ablation: top 12 + 6m kill, no near-high",
            {
                **ride_rockets_champ_defaults(top_n=12),
                "near_high_min_frac": 0.0,
                "selection_filters": ["ride_rockets", "champ", "kill_only", "top_12"],
            },
        ),
        (
            "champ_top15_near_only",
            "Ablation: top 15 + near-high only (no 6m kill)",
            {
                **ride_rockets_champ_defaults(top_n=15),
                "require_positive_raw_6": False,
                "selection_filters": ["ride_rockets", "champ", "near_only", "top_15"],
            },
        ),
    ]


def ride_rockets_variant_catalog() -> list[tuple[str, str, dict]]:
    """
    Ten creative ride-the-rockets variants.

    Returns list of (slug, thesis, config_patch overlapping ride_rockets_defaults).
    """
    base_n = 25
    return [
        (
            "01_core_monthly25",
            "Baseline: top 25, monthly, 70/30 blend, SMA150, score×cap, 15% cap",
            ride_rockets_defaults(top_n=base_n),
        ),
        (
            "02_ten_rockets",
            "Ultra-concentrated top 10 with 20% single-name ceiling",
            {
                **ride_rockets_defaults(top_n=10),
                "max_single_weight": 0.20,
                "score_weight_power": 1.5,
                "selection_filters": ["ride_rockets", "ten_rockets"],
            },
        ),
        (
            "03_acceleration",
            "Weight fresh thrust: 80% 6m / 20% 12m risk-adjusted momentum",
            {
                **ride_rockets_defaults(top_n=25),
                "mom_weight_6m": 0.8,
                "mom_weight_12m": 0.2,
                "selection_filters": ["ride_rockets", "acceleration"],
            },
        ),
        (
            "04_kill_on_6m_fade",
            "Exit when intermediate 6m return turns negative (thrust died)",
            {
                **ride_rockets_defaults(top_n=25),
                "require_positive_raw_6": True,
                "selection_filters": ["ride_rockets", "kill_6m"],
            },
        ),
        (
            "05_mega_only",
            "Only $40B+ mega-caps — NVDA/AVGO hunting ground",
            {
                **ride_rockets_defaults(top_n=20),
                "min_market_cap_usd": 40_000_000_000.0,
                "selection_filters": ["ride_rockets", "mega_only"],
            },
        ),
        (
            "06_near_52w_high",
            "Must sit within 5% of 52-week high (extension riders only)",
            {
                **ride_rockets_defaults(top_n=25),
                "near_high_min_frac": 0.95,
                "selection_filters": ["ride_rockets", "near_high"],
            },
        ),
        (
            "07_equal_weight20",
            "Equal-weight top 20 — ignore size, pure relative strength",
            {
                **ride_rockets_defaults(top_n=20),
                "cap_proxy": "equal",
                "selection_filters": ["ride_rockets", "equal_weight"],
            },
        ),
        (
            "08_raw_price_rockets",
            "Rank by raw price momentum (not vol-adjusted) — favor volatility rockets",
            {
                **ride_rockets_defaults(top_n=20),
                "use_raw_momentum_rank": True,
                "max_single_weight": 0.12,
                "selection_filters": ["ride_rockets", "raw_rank"],
            },
        ),
        (
            "09_sticky_buffer",
            "Top 25 with 1.5× hold buffer — ride longer, less churn",
            {
                **ride_rockets_defaults(top_n=25),
                "use_buffer_rule": True,
                "hold_buffer_mult": 1.5,
                "selection_filters": ["ride_rockets", "sticky_buffer"],
            },
        ),
        (
            "10_power_tilt15",
            "Top 15, score^2 tilt, 25% cap — pile into the strongest rockets",
            {
                **ride_rockets_defaults(top_n=15),
                "score_weight_power": 2.0,
                "max_single_weight": 0.25,
                "weight_by_score_only": True,
                "require_positive_raw_3": True,
                "selection_filters": ["ride_rockets", "power_tilt"],
            },
        ),
    ]


@dataclass
class Sp500MomentumIndex:
    """Long-only S&P 500 momentum index approximation (SPMO-style)."""

    config: Sp500MomentumConfig = Sp500MomentumConfig()
    sector_map: Dict[str, str] | None = None
    membership_mask: pd.DataFrame | None = None
    market_cap_df: pd.DataFrame | None = None
    spy_panel: pd.DataFrame | None = None

    def _volume_series(self, df: pd.DataFrame) -> pd.Series | None:
        for col in ("volume", "Volume"):
            if col in df.columns:
                return df[col].astype(np.float64)
        return None

    def _build_panels(
        self, equity_dict: Dict[str, pd.DataFrame]
    ) -> dict[str, pd.DataFrame]:
        sma_w = int(self.config.above_sma_window) if self.config.above_sma_window > 0 else 0
        panels: dict[str, list[pd.Series]] = {
            k: []
            for k in (
                "risk_12",
                "risk_6",
                "raw_12",
                "raw_6",
                "raw_3",
                "vol_12",
                "ret",
                "close",
                "td",
                "above_sma",
                "near_high",
                "adv",
            )
        }
        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
            vol = self._volume_series(df)
            if vol is not None:
                vol.index = idx
            feats = compute_momentum_features(close.sort_index(), vol, sma_window=sma_w)
            for key in (
                "risk_12",
                "risk_6",
                "raw_12",
                "raw_6",
                "raw_3",
                "vol_12",
                "ret",
                "trading_days",
                "above_sma",
                "near_high",
                "adv_usd",
            ):
                out_key = "td" if key == "trading_days" else ("adv" if key == "adv_usd" else key)
                panels[out_key].append(feats[key].rename(t))
            panels["close"].append(close.rename(t))

        if not panels["risk_12"]:
            raise ValueError("no valid tickers with close prices")

        out: dict[str, pd.DataFrame] = {}
        master = None
        for key, ser_list in panels.items():
            df, _ = _align_panel_frames(ser_list)
            if key in ("ret",):
                out[key] = df.sort_index()
            else:
                out[key] = df.sort_index().ffill()
            master = df.index if master is None else master.union(df.index)
        master = master.sort_values()
        for key in out:
            out[key] = out[key].reindex(master)
            if key != "ret":
                out[key] = out[key].ffill()
        return out

    def _build_spy_features(self, master_index: pd.DatetimeIndex) -> pd.DataFrame:
        if self.spy_panel is None or self.spy_panel.empty:
            return pd.DataFrame(index=master_index)
        sp = self.spy_panel.copy()
        sp.index = pd.to_datetime(sp.index).tz_localize(None)
        close = sp["close"].astype(np.float64)
        ret = sp["ret"].astype(np.float64) if "ret" in sp.columns else close.pct_change()
        raw_12 = close.shift(21) / close.shift(252) - 1.0
        beta_win = int(self.config.residual_beta_window)
        var_spy = ret.rolling(beta_win, min_periods=beta_win).var(ddof=1)
        feat = pd.DataFrame({"spy_raw_12": raw_12, "spy_ret": ret}, index=close.index)
        feat = feat.reindex(master_index).ffill()
        return feat

    def _build_beta_panel(self, ret_df: pd.DataFrame, spy_ret: pd.Series) -> pd.DataFrame:
        win = int(self.config.residual_beta_window)
        var_spy = spy_ret.rolling(win, min_periods=win).var(ddof=1).replace(0.0, np.nan)
        beta_pan = []
        for t in ret_df.columns:
            cov = ret_df[t].rolling(win, min_periods=win).cov(spy_ret)
            beta_pan.append((cov / var_spy).rename(t))
        beta_df, _ = _align_panel_frames(beta_pan)
        return beta_df.reindex(ret_df.index)

    def _rebalance_index(self, panel: pd.DataFrame) -> pd.DatetimeIndex:
        monthly = panel.resample("BME").last().index
        if self.config.rebalance == "monthly":
            return monthly
        months = set(int(m) for m in self.config.semiannual_months)
        return monthly[monthly.month.isin(months)]

    @staticmethod
    def _iterative_cap(weights: np.ndarray, caps: np.ndarray, max_iter: int = 30) -> np.ndarray:
        w = np.array(weights, dtype=np.float64, copy=True)
        caps = np.asarray(caps, dtype=np.float64)
        for _ in range(max_iter):
            over = w > caps + 1e-12
            if not over.any():
                break
            excess = float(w[over].sum() - caps[over].sum())
            w[over] = caps[over]
            under = ~over & (w > 0)
            us = float(w[under].sum())
            if us <= 1e-12:
                break
            w[under] += excess * (w[under] / us)
        s = float(w.sum())
        if s > 1e-12:
            w /= s
        return w

    def _cap_sector_weights(self, weights: np.ndarray, tickers: list[str], sector_cap: float) -> np.ndarray:
        if sector_cap <= 0 or self.sector_map is None:
            return weights
        w = np.array(weights, dtype=np.float64, copy=True)
        sectors = [str(self.sector_map.get(t, "Unknown")) for t in tickers]
        for _ in range(20):
            sec_sums: dict[str, float] = {}
            for wi, sec in zip(w, sectors):
                if wi > 1e-12:
                    sec_sums[sec] = sec_sums.get(sec, 0.0) + float(wi)
            over_secs = {s for s, v in sec_sums.items() if v > sector_cap + 1e-12}
            if not over_secs:
                break
            for sec in over_secs:
                idxs = [i for i, s in enumerate(sectors) if s == sec]
                sec_w = float(w[idxs].sum())
                if sec_w <= 1e-12:
                    continue
                w[idxs] *= sector_cap / sec_w
            s = float(w.sum())
            if s > 1e-12:
                w /= s
        return w

    def _target_count(self, n_eligible: int) -> int:
        cfg = self.config
        if cfg.top_n > 0:
            return min(int(cfg.top_n), n_eligible)
        if cfg.quintile_selection:
            return max(1, int(round(n_eligible / 5.0)))
        return min(100, n_eligible)

    def _compose_risk_row(
        self,
        risk_6: np.ndarray,
        risk_12: np.ndarray,
        raw_12: np.ndarray,
        vol_12: np.ndarray,
        beta: np.ndarray | None,
        spy_raw_12: float,
        member_row: np.ndarray | None,
        raw_6: np.ndarray | None = None,
    ) -> np.ndarray:
        cfg = self.config
        if cfg.use_raw_momentum_rank:
            if raw_6 is None:
                blended = raw_12.astype(np.float64)
            else:
                blended = blend_horizon_risk_adj(
                    raw_6, raw_12, weight_6m=cfg.mom_weight_6m, weight_12m=cfg.mom_weight_12m
                )
        else:
            blended = blend_horizon_risk_adj(
                risk_6, risk_12, weight_6m=cfg.mom_weight_6m, weight_12m=cfg.mom_weight_12m
            )
            if cfg.use_residual_momentum and beta is not None and np.isfinite(spy_raw_12):
                blended = residualize_risk_row(raw_12, vol_12, beta, spy_raw_12)
        risk_for_z = blended.astype(np.float64).copy()
        if member_row is not None:
            risk_for_z[~member_row.astype(bool)] = np.nan
        if cfg.use_sp_index_rules:
            return cross_sectional_sp_momentum_scores(risk_for_z)
        return risk_for_z

    def _weights_at_rebalance(
        self,
        panels_row: dict[str, np.ndarray],
        member_row: np.ndarray | None,
        prev_held_idx: set[int],
        tickers: list[str],
        spy_raw_12: float,
    ) -> np.ndarray:
        cfg = self.config
        n = len(tickers)
        out = np.zeros(n, dtype=np.float64)

        scores = self._compose_risk_row(
            panels_row["risk_6"],
            panels_row["risk_12"],
            panels_row["raw_12"],
            panels_row["vol_12"],
            panels_row.get("beta"),
            spy_raw_12,
            member_row,
            raw_6=panels_row.get("raw_6"),
        )

        eligible = np.isfinite(scores) & (scores > float(cfg.min_mom_score))
        if member_row is not None:
            eligible &= member_row.astype(bool)
        if cfg.min_trading_days > 0:
            eligible &= np.isfinite(panels_row["td"]) & (panels_row["td"] >= float(cfg.min_trading_days))
        if cfg.require_positive_raw_mom or cfg.min_raw_mom > 0.0:
            floor = float(cfg.min_raw_mom)
            eligible &= np.isfinite(panels_row["raw_12"]) & (panels_row["raw_12"] > floor)
        if cfg.require_positive_raw_6:
            eligible &= np.isfinite(panels_row["raw_6"]) & (panels_row["raw_6"] > 0.0)
        if cfg.require_positive_raw_3:
            eligible &= np.isfinite(panels_row["raw_3"]) & (panels_row["raw_3"] > 0.0)
        if cfg.above_sma_window > 0:
            eligible &= np.isfinite(panels_row["above_sma"]) & (panels_row["above_sma"] > 0.5)
        if cfg.near_high_min_frac > 0.0 and panels_row.get("near_high") is not None:
            eligible &= np.isfinite(panels_row["near_high"]) & (
                panels_row["near_high"] >= float(cfg.near_high_min_frac)
            )
        if cfg.min_market_cap_usd > 0.0 and panels_row.get("mcap") is not None:
            eligible &= np.isfinite(panels_row["mcap"]) & (panels_row["mcap"] >= float(cfg.min_market_cap_usd))
        if cfg.min_adv_usd > 0.0:
            eligible &= np.isfinite(panels_row["adv"]) & (panels_row["adv"] >= float(cfg.min_adv_usd))

        idx_valid = np.nonzero(eligible)[0]
        if idx_valid.size == 0:
            return out

        target = self._target_count(int(idx_valid.size))
        target = adjust_target_for_rank_gap(scores, eligible, target, float(cfg.min_rank_score_pct_gap))
        if target <= 0:
            return out

        buffer_mult = float(cfg.hold_buffer_mult) if float(cfg.hold_buffer_mult) > 1.0 else 1.2
        if cfg.sector_neutral_select and self.sector_map is not None:
            sector_ids = np.array([str(self.sector_map.get(tickers[j], "Unknown")) for j in range(n)], dtype=object)
            chosen_mask = select_sector_neutral_long(scores, eligible, sector_ids, target)
        elif cfg.use_buffer_rule and (cfg.use_sp_index_rules or float(cfg.hold_buffer_mult) > 1.0):
            chosen_mask = select_with_sp_buffer(
                scores, eligible, prev_held_idx, target, buffer_mult=buffer_mult
            )
        else:
            sub = np.where(eligible, scores, -np.inf)
            k = min(target, int(idx_valid.size))
            pick = np.argpartition(-sub, k - 1)[:k]
            chosen_mask = np.zeros(n, dtype=bool)
            chosen_mask[pick] = True
            chosen_mask &= eligible

        chosen = np.nonzero(chosen_mask)[0]
        if chosen.size == 0:
            return out

        close_row = panels_row["close"]
        mcap_row = panels_row.get("mcap")

        if cfg.cap_proxy == "equal":
            out[chosen] = 1.0 / float(chosen.size)
            return out

        if cfg.weight_by_score_only:
            score_vals = np.maximum(scores[chosen], 1e-8) ** float(cfg.score_weight_power)
            denom = float(score_vals.sum())
            if denom <= 1e-12:
                out[chosen] = 1.0 / float(chosen.size)
                return out
            w = np.zeros(n, dtype=np.float64)
            w[chosen] = score_vals / denom
            if cfg.max_single_weight > 0.0:
                caps = np.full(chosen.size, float(cfg.max_single_weight))
                w[chosen] = self._iterative_cap(w[chosen], caps)
            return w

        if cfg.cap_proxy == "mkt_cap" and mcap_row is not None:
            cap_vals = mcap_row[chosen]
        else:
            cap_vals = close_row[chosen]

        score_vals = np.maximum(scores[chosen], 1e-8) ** float(cfg.score_weight_power)
        numer = np.where(np.isfinite(cap_vals) & (cap_vals > 0.0), cap_vals * score_vals, 0.0)
        denom = float(numer.sum())
        if denom <= 1e-12:
            out[chosen] = 1.0 / float(chosen.size)
            return out

        w = np.zeros(n, dtype=np.float64)
        w[chosen] = numer / denom

        if cfg.use_sp_index_rules and mcap_row is not None and np.isfinite(mcap_row).any():
            uni_mcap = np.where(member_row if member_row is not None else True, mcap_row, 0.0)
            uni_mcap = np.where(np.isfinite(uni_mcap) & (uni_mcap > 0), uni_mcap, 0.0)
            uni_total = float(uni_mcap.sum())
            if uni_total > 1e-12:
                uni_w = uni_mcap / uni_total
                caps = np.minimum(float(cfg.max_single_weight), float(cfg.mkt_cap_multiplier) * uni_w)
            else:
                caps = np.full(n, float(cfg.max_single_weight))
        elif cfg.max_single_weight > 0.0:
            caps = np.full(n, float(cfg.max_single_weight))
        else:
            caps = np.ones(n)

        w[chosen] = self._iterative_cap(w[chosen], caps[chosen])
        if cfg.max_sector_weight > 0.0:
            w = self._cap_sector_weights(w, tickers, float(cfg.max_sector_weight))
        return w

    def _rebalance_rows(
        self,
        panels: dict[str, pd.DataFrame],
        bm_index: pd.DatetimeIndex,
        tickers: list[str],
    ) -> np.ndarray:
        mem_df = self.membership_mask
        if mem_df is not None:
            mem_df = mem_df.reindex(bm_index).reindex(columns=tickers, fill_value=False)

        mcap_m = None
        if self.market_cap_df is not None:
            mcap_m = self.market_cap_df.reindex(bm_index, method="ffill").reindex(columns=tickers)
        elif self.config.cap_proxy == "mkt_cap":
            mcap_m = panels["close"].reindex(bm_index, method="ffill")

        spy_feat = self._build_spy_features(panels["close"].index)
        beta_df = None
        if self.config.use_residual_momentum and "spy_ret" in spy_feat.columns:
            beta_df = self._build_beta_panel(panels["ret"], spy_feat["spy_ret"])

        keys = ("risk_6", "risk_12", "raw_12", "raw_6", "raw_3", "vol_12", "close", "td", "above_sma", "near_high", "adv")
        at_bm = {k: panels[k].reindex(bm_index, method="ffill") for k in keys if k in panels}
        if mcap_m is not None:
            at_bm["mcap"] = mcap_m
        if beta_df is not None:
            at_bm["beta"] = beta_df.reindex(bm_index, method="ffill")

        m, n = len(bm_index), len(tickers)
        weights_m = np.zeros((m, n), dtype=np.float64)
        prev_held: set[int] = set()

        for i, sig_date in enumerate(bm_index):
            row = {k: at_bm[k].iloc[i].to_numpy(dtype=np.float64) for k in at_bm}
            mem_row = mem_df.iloc[i].to_numpy(dtype=bool) if mem_df is not None else None
            spy_raw = float(spy_feat["spy_raw_12"].reindex(bm_index).iloc[i]) if "spy_raw_12" in spy_feat.columns else float("nan")
            w = self._weights_at_rebalance(row, mem_row, prev_held, tickers, spy_raw)
            weights_m[i] = w
            prev_held = {j for j in range(n) if w[j] > 1e-8}
        return weights_m

    def generate_returns(self, equity_dict: Dict[str, pd.DataFrame], *, verbose: bool = True) -> pd.Series:
        panels = self._build_panels(equity_dict)
        bm_index = self._rebalance_index(panels["risk_12"])
        tickers = list(panels["risk_12"].columns)
        weights_m = self._rebalance_rows(panels, bm_index, tickers)

        master_index = panels["ret"].index.sort_values()
        weights_m_df = pd.DataFrame(weights_m, index=bm_index, columns=tickers)
        weights_d = weights_m_df.reindex(master_index).ffill().shift(1).fillna(0.0)

        ret_filled = panels["ret"].reindex(master_index).fillna(0.0).astype(np.float64)
        daily_rf = float(self.config.cash_annual_yield) / 252.0
        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 and len(bm_index) > 0:
            avg_names = float(np.mean([np.count_nonzero(weights_m[i] > 0) for i in range(len(bm_index))]))
            avg_wsum = float(np.mean(weights_m.sum(axis=1)))
            filt = ",".join(self.config.selection_filters) if self.config.selection_filters else "base"
            print(
                f"  SP500 momentum [{filt}]: target≈{self._target_count(len(tickers))} | "
                f"avg held/rebal ≈ {avg_names:.1f} | avg gross ≈ {avg_wsum:.2f}",
                flush=True,
            )
        return pd.Series(out, index=master_index, name="sp500_mom_ret", dtype=np.float64)

    def generate_rebalance_log(self, equity_dict: Dict[str, pd.DataFrame]) -> pd.DataFrame:
        panels = self._build_panels(equity_dict)
        bm_index = self._rebalance_index(panels["risk_12"])
        tickers = list(panels["risk_12"].columns)
        weights_m = self._rebalance_rows(panels, bm_index, tickers)
        master_index = panels["risk_12"].index.sort_values()

        mem_df = self.membership_mask
        if mem_df is not None:
            mem_df = mem_df.reindex(bm_index).reindex(columns=tickers, fill_value=False)
        spy_feat = self._build_spy_features(panels["close"].index)
        beta_df = None
        if self.config.use_residual_momentum:
            beta_df = self._build_beta_panel(panels["ret"], spy_feat["spy_ret"])

        keys = ("risk_6", "risk_12", "raw_12", "raw_6", "raw_3", "vol_12", "close", "td", "above_sma", "near_high", "adv")
        at_bm = {k: panels[k].reindex(bm_index, method="ffill") for k in keys if k in panels}
        if self.market_cap_df is not None:
            at_bm["mcap"] = self.market_cap_df.reindex(bm_index, method="ffill").reindex(columns=tickers)
        if beta_df is not None:
            at_bm["beta"] = beta_df.reindex(bm_index, method="ffill")

        rows: list[dict] = []
        prev_held: set[str] = set()
        prev_held_idx: set[int] = set()

        for i, sig_date in enumerate(bm_index):
            row = {k: at_bm[k].iloc[i].to_numpy(dtype=np.float64) for k in at_bm}
            mem_row = mem_df.iloc[i].to_numpy(dtype=bool) if mem_df is not None else None
            spy_raw = float(spy_feat["spy_raw_12"].reindex(bm_index).iloc[i])
            scores = self._compose_risk_row(
                row["risk_6"], row["risk_12"], row["raw_12"], row["vol_12"],
                row.get("beta"), spy_raw, mem_row, raw_6=row.get("raw_6"),
            )
            w = self._weights_at_rebalance(row, mem_row, prev_held_idx, tickers, spy_raw)
            held = {tickers[j] for j in range(len(tickers)) if w[j] > 1e-8}
            future = master_index[master_index > sig_date]
            effective = future.min() if len(future) else sig_date + pd.Timedelta(days=1)

            ranked = sorted(
                ((tickers[j], float(scores[j]), float(w[j])) for j in range(len(tickers)) if w[j] > 1e-8),
                key=lambda x: -x[1],
            )
            for rank, (ticker, mom_score, weight) in enumerate(ranked, start=1):
                rows.append({
                    "signal_date": sig_date.strftime("%Y-%m-%d"),
                    "effective_date": pd.Timestamp(effective).strftime("%Y-%m-%d"),
                    "ticker": ticker,
                    "sector": str(self.sector_map.get(ticker, "")) if self.sector_map else "",
                    "weight": round(weight, 6),
                    "mom_score": round(mom_score, 6),
                    "raw_mom_12": round(float(row["raw_12"][tickers.index(ticker)]), 6),
                    "rank": rank,
                    "is_new_entry": ticker not in prev_held,
                    "is_exit_from_prior": False,
                    "top_n": int(self._target_count(len(tickers))),
                })
            for t in prev_held - held:
                rows.append({
                    "signal_date": sig_date.strftime("%Y-%m-%d"),
                    "effective_date": pd.Timestamp(effective).strftime("%Y-%m-%d"),
                    "ticker": t,
                    "sector": str(self.sector_map.get(t, "")) if self.sector_map else "",
                    "weight": 0.0,
                    "mom_score": float("nan"),
                    "raw_mom_12": float("nan"),
                    "rank": 0,
                    "is_new_entry": False,
                    "is_exit_from_prior": True,
                    "top_n": int(self._target_count(len(tickers))),
                })
            prev_held = held
            prev_held_idx = {j for j in range(len(tickers)) if tickers[j] in held}
        return pd.DataFrame(rows)
