"""
ML-driven cross-sectional momentum: XGBoost predicts 21-day forward returns on the S&P 500
universe, with walk-forward training (no lookahead).

Features include stock-level signals plus **market context**: rolling beta / correlation to SPY,
momentum relative to SPY, and VIX level / stress vs a 20-day mean (via ^VIX).

By default the model is trained to predict **excess** forward return (stock minus SPY over the same
``horizon`` bars on each ticker's calendar); set ``target_mode="total"`` for raw stock forward return.

Requires: pandas, numpy, xgboost; :mod:`RenTech.strategy_stack.data_loader` (yfinance) for SPY/^VIX.
"""

from __future__ import annotations

import urllib.request
from dataclasses import dataclass, field
from io import StringIO
from typing import Any, Dict, Iterable, Literal

import numpy as np
import pandas as pd

_USER_AGENT = (
    "Mozilla/5.0 (compatible; ml_momentum_engine/1.0) "
    "AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"
)

try:
    from xgboost import XGBRegressor
except ImportError as e:  # pragma: no cover
    XGBRegressor = None  # type: ignore[misc, assignment]
    _XGB_IMPORT_ERROR = e
else:
    _XGB_IMPORT_ERROR = None


def _require_xgboost() -> None:
    if XGBRegressor is None:
        raise ImportError(
            "xgboost is required for MLMomentumEngine. Install with: pip install xgboost\n"
            f"Original error: {_XGB_IMPORT_ERROR}"
        )


def fetch_sp500_constituents_with_sectors() -> pd.DataFrame:
    """
    Download current S&P 500 constituents and GICS sectors from Wikipedia.

    Returns
    -------
    pd.DataFrame
        Columns: ``ticker`` (Yahoo-style), ``sector``.
    """
    url = "https://en.wikipedia.org/wiki/List_of_S%26P_500_companies"
    req = urllib.request.Request(url, headers={"User-Agent": _USER_AGENT})
    with urllib.request.urlopen(req, timeout=60) as resp:
        html = resp.read().decode("utf-8", errors="replace")
    try:
        tables = pd.read_html(StringIO(html))
    except ImportError as e:
        raise ImportError(
            "pandas.read_html needs an HTML parser. Install: pip install lxml html5lib beautifulsoup4"
        ) from e
    df = tables[0]
    sym_col = "Symbol" if "Symbol" in df.columns else next(c for c in df.columns if "symbol" in str(c).lower())
    sec_col = next(
        (c for c in df.columns if "gics" in str(c).lower() and "sector" in str(c).lower()),
        "GICS Sector" if "GICS Sector" in df.columns else None,
    )
    if sec_col is None:
        raise RuntimeError("Could not find GICS Sector column in S&P 500 Wikipedia table.")
    out = pd.DataFrame(
        {
            "ticker": df[sym_col].astype(str).str.replace(".", "-", regex=False).str.strip(),
            "sector": df[sec_col].astype(str).str.strip(),
        }
    )
    return out.drop_duplicates(subset=["ticker"]).reset_index(drop=True)


def _to_yahoo_symbol(symbol: str) -> str:
    return symbol.strip().upper().replace(".", "-")


def _union_index(frames: Iterable[pd.DataFrame]) -> pd.DatetimeIndex:
    idx: pd.DatetimeIndex | None = None
    for f in frames:
        if f is None or f.empty:
            continue
        ii = pd.to_datetime(f.index).tz_localize(None)
        idx = ii if idx is None else idx.union(ii)
    if idx is None:
        raise ValueError("No non-empty frames")
    return idx.sort_values()


def _series_ffill_on_calendar(s: pd.Series, cal: pd.DatetimeIndex) -> pd.Series:
    """Align ``s`` to ``cal`` using forward-fill along a merged timeline (point-in-time)."""
    s = s.dropna().astype(np.float64).sort_index()
    cal = pd.DatetimeIndex(cal).sort_values()
    if s.empty:
        return pd.Series(np.nan, index=cal, dtype=np.float64)
    u = cal.union(s.index).sort_values()
    return s.reindex(u).ffill().reindex(cal)


def build_market_panel(
    master_index: pd.DatetimeIndex,
    *,
    fetch_period: str = "10y",
) -> pd.DataFrame:
    """
    SPY + VIX features on a calendar that covers ``master_index`` and index history for betas.

    Columns include SPY returns / momentum and VIX level / stress vs 20d mean.
    """
    from RenTech.strategy_stack.data_loader import DataLoader

    dl = DataLoader()
    spy_df = dl.fetch_daily("SPY", period=fetch_period)
    vix_df = dl.fetch_daily("^VIX", period=fetch_period)
    if spy_df.empty or "close" not in spy_df.columns:
        raise RuntimeError("Could not load SPY daily data for market context.")
    if vix_df.empty or "close" not in vix_df.columns:
        raise RuntimeError("Could not load ^VIX daily data for market context.")

    mi = pd.DatetimeIndex(master_index).sort_values()
    spy_idx = pd.to_datetime(spy_df.index).tz_localize(None)
    vix_idx = pd.to_datetime(vix_df.index).tz_localize(None)
    cal = mi.union(spy_idx).union(vix_idx).sort_values()

    spy_c = spy_df["close"].astype(np.float64)
    spy_c.index = spy_idx
    vix_c = vix_df["close"].astype(np.float64)
    vix_c.index = vix_idx

    spy_al = _series_ffill_on_calendar(spy_c, cal)
    vix_al = _series_ffill_on_calendar(vix_c, cal)

    out = pd.DataFrame(index=cal)
    out["spy_close"] = spy_al
    out["spy_ret"] = out["spy_close"].pct_change()
    out["spy_mom_21"] = out["spy_close"] / out["spy_close"].shift(21) - 1.0
    out["spy_mom_63"] = out["spy_close"] / out["spy_close"].shift(63) - 1.0
    out["spy_mom_252"] = out["spy_close"] / out["spy_close"].shift(252) - 1.0
    out["vix"] = vix_al
    vix_ma20 = out["vix"].rolling(20, min_periods=20).mean()
    out["vix_ratio_ma20"] = out["vix"] / vix_ma20.replace(0.0, np.nan) - 1.0
    out["vix_chg_20"] = out["vix"] / out["vix"].shift(20).replace(0.0, np.nan) - 1.0
    spy_vol_20 = out["spy_ret"].rolling(20, min_periods=20).std(ddof=1)
    out["spy_vol_20"] = spy_vol_20
    return out


def _compute_ticker_features_and_target(
    df: pd.DataFrame,
    market: pd.DataFrame,
    *,
    horizon: int = 21,
    target_mode: Literal["excess", "total"] = "excess",
) -> tuple[pd.DataFrame, pd.Series]:
    """
    Point-in-time features at close t.

    Target (``target_mode == "excess"``): forward ``horizon``-bar return of the stock minus the
    same-horizon forward return of SPY, both measured on **this ticker's** date index (aligned
    ``spy_close`` from ``market``). Known only after ``t + horizon`` bars.

    Target (``target_mode == "total"``): stock forward total return only.
    """
    need = {"close", "high", "volume"}
    miss = need.difference(df.columns)
    if miss:
        raise KeyError(f"DataFrame missing columns {sorted(miss)}; need close, high, volume.")

    df = df.sort_index()
    m = market.reindex(df.index)

    c = df["close"].astype(np.float64)
    h = df["high"].astype(np.float64)
    v = df["volume"].astype(np.float64)

    feat = pd.DataFrame(index=df.index)
    feat["mom_1m"] = c / c.shift(21) - 1.0
    feat["mom_3m"] = c / c.shift(63) - 1.0
    feat["mom_6m"] = c / c.shift(126) - 1.0
    feat["mom_12m"] = c / c.shift(252) - 1.0
    r = c.pct_change()
    feat["vol_20"] = r.rolling(window=20, min_periods=20).std(ddof=1)
    feat["vol_60"] = r.rolling(window=60, min_periods=60).std(ddof=1)
    v5 = v.rolling(window=5, min_periods=5).mean()
    v20 = v.rolling(window=20, min_periods=20).mean()
    feat["vol_spike"] = v5 / v20.replace(0.0, np.nan) - 1.0
    roll_max = c.rolling(window=252, min_periods=252).max()
    feat["dist_high_52w"] = c / roll_max.replace(0.0, np.nan) - 1.0

    spy_r = m["spy_ret"].astype(np.float64)
    var_spy = spy_r.rolling(60, min_periods=60).var(ddof=1).replace(0.0, np.nan)
    cov_rs = r.rolling(60, min_periods=60).cov(spy_r)
    feat["beta_60"] = cov_rs / var_spy
    feat["corr_60"] = r.rolling(60, min_periods=60).corr(spy_r)
    feat["rel_mom_21"] = feat["mom_1m"] - m["spy_mom_21"]
    feat["rel_mom_63"] = feat["mom_3m"] - m["spy_mom_63"]
    feat["rel_mom_252"] = feat["mom_12m"] - m["spy_mom_252"]
    spy_v20 = m["spy_vol_20"].astype(np.float64)
    feat["vol_ratio_vs_spy_20"] = feat["vol_20"] / spy_v20.replace(0.0, np.nan)
    feat["mkt_vix"] = m["vix"]
    feat["mkt_vix_ratio_ma20"] = m["vix_ratio_ma20"]
    feat["mkt_vix_chg_20"] = m["vix_chg_20"]
    feat["mkt_spy_mom_63"] = m["spy_mom_63"]
    feat["mkt_spy_mom_252"] = m["spy_mom_252"]

    stock_fwd = c.shift(-horizon) / c - 1.0
    if target_mode == "total":
        target = stock_fwd
    else:
        spy_c = m["spy_close"].astype(np.float64)
        spy_fwd = spy_c.shift(-horizon) / spy_c - 1.0
        target = stock_fwd - spy_fwd
    target.name = "target_fwd_excess" if target_mode == "excess" else "target_fwd_total"
    return feat, target


def _one_hot_sector_matrix(
    tickers: list[str],
    sector_map: Dict[str, str],
    sector_universe: list[str],
) -> pd.DataFrame:
    """Rows aligned to ``tickers``; columns sector_<name>."""
    cols = [f"sector_{_sanitize(s)}" for s in sector_universe]
    out = pd.DataFrame(0.0, index=np.arange(len(tickers)), columns=cols, dtype=np.float64)
    sec_to_col = {s: f"sector_{_sanitize(s)}" for s in sector_universe}
    unk = sec_to_col.get("Unknown")
    for i, t in enumerate(tickers):
        s = sector_map.get(t, "Unknown")
        col = sec_to_col.get(s)
        if col is None or col not in out.columns:
            col = unk
        if col is not None and col in out.columns:
            out.iloc[i, out.columns.get_loc(col)] = 1.0
    return out


def _sanitize(s: str) -> str:
    return "".join(ch if ch.isalnum() else "_" for ch in str(s))[:48]


def _stack_training_rows(
    tickers: list[str],
    train_dates: pd.DatetimeIndex,
    feat_by_ticker: Dict[str, pd.DataFrame],
    target_by_ticker: Dict[str, pd.Series],
    sector_map: Dict[str, str],
    sector_universe: list[str],
) -> tuple[pd.DataFrame, np.ndarray]:
    """
    Stack (date, ticker) rows for model training.

    Crucial point-in-time hygiene:
      - do NOT reindex to a union calendar and fill missing with 0
      - only include rows where all numeric features AND the target are finite
    """
    rows_x: list[pd.DataFrame] = []
    rows_y: list[np.ndarray] = []
    sec_cols = [f"sector_{_sanitize(s)}" for s in sector_universe]

    for t in tickers:
        if t not in feat_by_ticker or t not in target_by_ticker:
            continue

        Xf = feat_by_ticker[t].reindex(train_dates)
        yv = target_by_ticker[t].reindex(train_dates).astype(np.float64)

        # Constant sector one-hot for this ticker.
        oh = _one_hot_sector_matrix([t], sector_map, sector_universe)
        oh_vals = np.tile(oh.iloc[0].values.astype(np.float64), (len(train_dates), 1))
        oh_df = pd.DataFrame(oh_vals, index=train_dates, columns=sec_cols, dtype=np.float64)

        block = pd.concat([Xf, oh_df], axis=1)

        # Keep only rows with fully-defined numeric features and target.
        numeric = block.drop(columns=sec_cols, errors="ignore")
        ok = numeric.notna().all(axis=1) & yv.notna()
        if not bool(ok.any()):
            continue

        block_ok = block.loc[ok].astype(np.float64)
        y_ok = yv.loc[ok].to_numpy(dtype=np.float64)
        rows_x.append(block_ok)
        rows_y.append(y_ok)

    if not rows_x:
        return pd.DataFrame(), np.array([])

    X = pd.concat(rows_x, axis=0).replace([np.inf, -np.inf], np.nan)
    y = np.concatenate(rows_y)
    finite = np.isfinite(y)
    if not bool(finite.any()):
        return pd.DataFrame(), np.array([])

    X = X.iloc[finite].dropna(axis=0, how="any")
    y = y[finite]
    if X.empty or len(y) < 10:
        return pd.DataFrame(), np.array([])
    return X, y


def _winsorize_fit(X: pd.DataFrame, cols: list[str], *, q_low: float = 0.01, q_high: float = 0.99) -> pd.DataFrame:
    qs = X[cols].quantile([q_low, q_high])
    lo = qs.loc[q_low]
    hi = qs.loc[q_high]
    out = X.copy()
    out[cols] = out[cols].clip(lower=lo, upper=hi, axis=1)
    return out


def _winsorize_apply(
    X: pd.DataFrame,
    cols: list[str],
    lo: pd.Series,
    hi: pd.Series,
) -> pd.DataFrame:
    out = X.copy()
    out[cols] = out[cols].clip(lower=lo, upper=hi, axis=1)
    return out


def _leg_risk_weights(names: list[str], vol_series: pd.Series | None, *, equal_fallback: bool) -> np.ndarray:
    """
    Within-leg weights that sum to 1.0: inverse-volatility if ``vol_series`` has valid entries,
    else equal weight.
    """
    n = len(names)
    if n == 0:
        return np.array([], dtype=np.float64)
    if equal_fallback or vol_series is None:
        return np.ones(n, dtype=np.float64) / float(n)
    v = vol_series.reindex(names).astype(np.float64).to_numpy()
    inv = np.where(np.isfinite(v) & (v > 1e-12), 1.0 / v, 0.0)
    ssum = float(inv.sum())
    if ssum <= 1e-18:
        return np.ones(n, dtype=np.float64) / float(n)
    return inv / ssum


def _build_monthly_weights_from_predictions(
    pred: pd.Series,
    all_tickers: list[str],
    top_n: int,
    *,
    vol_series: pd.Series | None = None,
    weighting: Literal["inv_vol", "equal"] = "inv_vol",
) -> pd.Series:
    """
    Long top_n predictions, short bottom_n. Each leg sums to +1 / −1 in absolute notionally.

    ``inv_vol``: risk-parity style within each leg using ``vol_series`` (e.g. 20d return vol at rebalance).
    ``equal``: legacy 1/top_n per name.
    """
    s = pred.astype(np.float64)
    valid = s.dropna()
    if len(valid) < 2 * top_n:
        return pd.Series(0.0, index=all_tickers, dtype=np.float64)

    longs = valid.nlargest(top_n).index.tolist()
    shorts = valid.nsmallest(top_n).index.tolist()
    w = pd.Series(0.0, index=all_tickers, dtype=np.float64)

    use_eq = weighting == "equal"
    lw = _leg_risk_weights(longs, vol_series, equal_fallback=use_eq)
    sw = _leg_risk_weights(shorts, vol_series, equal_fallback=use_eq)

    for i, t in enumerate(longs):
        w.loc[t] = float(lw[i])
    for i, t in enumerate(shorts):
        w.loc[t] = -float(sw[i])
    return w


@dataclass
class MLMomentumEngine:
    """
    Walk-forward XGBoost regressor on pooled (date, ticker) panels; monthly rebalance.

    Training uses only rows with ``date <= R - horizon`` so the forward return label is realized
    on or before rebalance day ``R`` (point-in-time).
    """

    train_years: int = 2
    horizon: int = 21
    top_n: int = 25
    # Yahoo period for SPY/^VIX (match equity horizon, e.g. 20y).
    data_period: str = "10y"
    # Train label: excess vs SPY over ``horizon`` bars, or total stock forward return.
    target_mode: Literal["excess", "total"] = "excess"
    # Within-leg sizing: inverse 20d vol (risk parity) vs equal weight.
    weighting: Literal["inv_vol", "equal"] = "inv_vol"
    xgb_params: dict[str, Any] = field(
        default_factory=lambda: {
            "max_depth": 3,
            "learning_rate": 0.01,
            "n_estimators": 100,
            "subsample": 0.8,
            "colsample_bytree": 0.8,
            "random_state": 42,
            "n_jobs": -1,
        }
    )

    def generate_ls_returns(
        self,
        equity_dict: Dict[str, pd.DataFrame],
        *,
        sector_map: Dict[str, str] | None = None,
        cash_annual_yield: float = 0.04,
        data_period: str | None = None,
        target_mode: Literal["excess", "total"] | None = None,
        weighting: Literal["inv_vol", "equal"] | None = None,
        verbose: bool = True,
    ) -> pd.Series:
        """
        Build daily L/S returns from ML predictions (monthly rebalance, next-bar execution).

        Parameters
        ----------
        equity_dict
            ticker -> daily OHLCV with at least ``close``, ``high``, ``volume`` (and index dates).
        sector_map
            ticker -> GICS sector name. If None, missing tickers use ``"Unknown"``.
        data_period
            Yahoo ``period`` for SPY and ^VIX loads (defaults to :attr:`data_period` on the engine).
        target_mode
            ``"excess"`` (default): predict stock minus SPY forward return; ``"total"``: stock only.
        weighting
            ``"inv_vol"`` (default): scale longs/shorts by inverse ``vol_20`` at rebalance; ``"equal"``: 1/N.
        cash_annual_yield
            Daily cash rate ``cash_annual_yield / 252`` applied to any capital not in the book
            when gross exposure differs from 100% long + 100% short (here weights sum to 0 net;
            gross = 2.0 when full legs).
        verbose
            Print walk-forward progress.

        Returns
        -------
        pd.Series
            Daily simple returns, indexed like the union calendar of inputs.
        """
        _require_xgboost()
        if not equity_dict:
            raise ValueError("equity_dict cannot be empty")

        tickers = sorted(equity_dict.keys())
        sector_map = dict(sector_map or {})
        for t in tickers:
            sector_map.setdefault(t, "Unknown")

        all_sectors = sorted(set(sector_map.values()))
        if "Unknown" not in all_sectors:
            all_sectors.append("Unknown")

        master_index = _union_index(equity_dict[t] for t in tickers)
        period_use = str(data_period or self.data_period).strip() or "10y"
        tm: Literal["excess", "total"] = self.target_mode if target_mode is None else target_mode
        if tm not in ("excess", "total"):
            raise ValueError(f"target_mode must be 'excess' or 'total', got {tm!r}")
        wt: Literal["inv_vol", "equal"] = self.weighting if weighting is None else weighting
        if wt not in ("inv_vol", "equal"):
            raise ValueError(f"weighting must be 'inv_vol' or 'equal', got {wt!r}")
        if verbose:
            print(f"  ML market context: loading SPY + ^VIX (period={period_use}) …")
            if tm == "excess":
                print(f"  ML target: excess ({self.horizon}-bar fwd stock minus SPY, ticker calendar)")
            else:
                print(f"  ML target: total ({self.horizon}-bar fwd stock return)")
            print(f"  ML leg sizing: {wt}" + (" (inverse 20d vol within long/short)" if wt == "inv_vol" else " (equal 1/N)"))
        market_wide = build_market_panel(master_index, fetch_period=period_use)

        # Per-ticker features & targets (target uses future prices — only consumed in training rows with t <= R - horizon)
        feat_by_ticker: Dict[str, pd.DataFrame] = {}
        target_by_ticker: Dict[str, pd.Series] = {}
        for t in tickers:
            df = equity_dict[t].copy()
            df.index = pd.to_datetime(df.index).tz_localize(None)
            df = df.sort_index()
            feat, tgt = _compute_ticker_features_and_target(
                df,
                market_wide,
                horizon=self.horizon,
                target_mode=tm,
            )
            feat_by_ticker[t] = feat
            target_by_ticker[t] = tgt

        # Warm-up: need 252d for 12m mom + 60d beta + label lag
        warmup = 320
        if len(master_index) <= warmup + self.horizon + 5:
            return pd.Series(0.0, index=master_index, name="ls_ret", dtype=np.float64)

        # Full trading calendar (daily rows for walk-forward training)
        mi = master_index.sort_values()
        # Month-end rebalance dates on the master calendar
        month_ends = mi.to_series().resample("BME").last().dropna()
        rebalance_dates = pd.DatetimeIndex(month_ends.values)
        rebalance_dates = rebalance_dates[rebalance_dates.isin(mi)]

        sec_cols = [f"sector_{_sanitize(s)}" for s in all_sectors]
        numeric_cols = [
            "mom_1m",
            "mom_3m",
            "mom_6m",
            "mom_12m",
            "vol_20",
            "vol_60",
            "vol_spike",
            "dist_high_52w",
            "beta_60",
            "corr_60",
            "rel_mom_21",
            "rel_mom_63",
            "rel_mom_252",
            "vol_ratio_vs_spy_20",
            "mkt_vix",
            "mkt_vix_ratio_ma20",
            "mkt_vix_chg_20",
            "mkt_spy_mom_63",
            "mkt_spy_mom_252",
        ]
        feature_cols = numeric_cols + sec_cols

        weights_rows: list[pd.Series] = []
        weight_index: list[pd.Timestamp] = []

        from pandas.tseries.offsets import BDay

        train_delta = pd.DateOffset(years=int(self.train_years))
        # Validation: last 3 months (~63 trading days) inside the training window.
        val_days = 63

        for R in rebalance_dates:
            # Last training date t with t + horizon trading days <= R  =>  t <= R - horizon BDays
            t_max = R - BDay(self.horizon)
            train_end = mi[mi <= t_max]
            if len(train_end) == 0:
                continue
            t_max_eff = train_end.max()

            t_min = t_max_eff - train_delta
            # Daily training rows (point-in-time): 2y window ending when 21d label is known by R
            train_dates = mi[(mi >= t_min) & (mi <= t_max_eff)]
            if len(train_dates) < 80:
                continue

            X_train, y_train = _stack_training_rows(
                tickers,
                train_dates,
                feat_by_ticker,
                target_by_ticker,
                sector_map,
                all_sectors,
            )
            if X_train.empty or len(y_train) < 2000:
                continue

            X_train = X_train.reindex(columns=feature_cols).astype(np.float64)

            # Time-based validation split by date index (kept in X_train).
            td = train_dates.sort_values()
            if len(td) <= val_days + 20:
                continue
            val_dates = set(td[-val_days:])
            is_val = X_train.index.isin(val_dates)
            mask_val = np.asarray(is_val, dtype=bool)
            mask_tr = ~mask_val
            # Use .iloc to avoid any ambiguity with duplicate date indices.
            X_tr = X_train.iloc[mask_tr]
            y_tr = y_train[mask_tr]
            X_va = X_train.iloc[mask_val]
            y_va = y_train[mask_val]
            if len(y_tr) < 1000 or len(y_va) < 200:
                continue

            # Winsorize numeric features based on training split (robustify outliers).
            qs = X_tr[numeric_cols].quantile([0.01, 0.99])
            lo = qs.loc[0.01]
            hi = qs.loc[0.99]
            X_tr = _winsorize_apply(X_tr, numeric_cols, lo, hi)
            X_va = _winsorize_apply(X_va, numeric_cols, lo, hi)

            model = XGBRegressor(
                **self.xgb_params,
                objective="reg:squarederror",
                eval_metric="rmse",
            )
            model.fit(
                X_tr.to_numpy(dtype=np.float64),
                y_tr,
                eval_set=[(X_va.to_numpy(dtype=np.float64), y_va)],
                verbose=False,
            )

            # Prediction features at R (point-in-time: known at close R)
            pred_rows: list[float] = []
            pred_tickers: list[str] = []
            for t in tickers:
                if R not in feat_by_ticker[t].index:
                    continue
                row = feat_by_ticker[t].loc[[R]].astype(np.float64)
                if row[numeric_cols].isna().any(axis=None):
                    continue
                oh = _one_hot_sector_matrix([t], sector_map, all_sectors)
                oh_arr = oh.iloc[0].reindex(sec_cols).fillna(0.0).to_numpy(dtype=np.float64)
                oh_df = pd.DataFrame(oh_arr.reshape(1, -1), index=[R], columns=sec_cols)
                Xp = pd.concat([row, oh_df], axis=1)
                Xp = Xp.reindex(columns=feature_cols, fill_value=0.0)
                Xp = _winsorize_apply(Xp, numeric_cols, lo, hi)
                p = float(model.predict(Xp.to_numpy(dtype=np.float64))[0])
                pred_rows.append(p)
                pred_tickers.append(t)

            if len(pred_rows) < 2 * self.top_n:
                continue

            pred_s = pd.Series(pred_rows, index=pred_tickers, dtype=np.float64)
            vol_map: dict[str, float] = {}
            if wt == "inv_vol":
                for tk in pred_tickers:
                    if tk not in feat_by_ticker or R not in feat_by_ticker[tk].index:
                        continue
                    vv = feat_by_ticker[tk].loc[R, "vol_20"]
                    if np.isfinite(vv) and float(vv) > 0:
                        vol_map[tk] = float(vv)
            vol_ser = pd.Series(vol_map, dtype=np.float64) if vol_map else None
            w = _build_monthly_weights_from_predictions(
                pred_s,
                tickers,
                self.top_n,
                vol_series=vol_ser,
                weighting=wt,
            )
            weights_rows.append(w.reindex(tickers).fillna(0.0))
            weight_index.append(R)

        if not weights_rows:
            return pd.Series(0.0, index=master_index, name="ls_ret", dtype=np.float64)

        weights_m_df = pd.DataFrame(weights_rows, index=pd.DatetimeIndex(weight_index), columns=tickers)
        weights_d = weights_m_df.reindex(master_index).ffill()
        weights_d = weights_d.shift(1).fillna(0.0)

        ret_df = pd.DataFrame(
            {t: equity_dict[t]["close"].astype(np.float64).pct_change().reindex(master_index) for t in tickers}
        )
        ret_filled = ret_df.fillna(0.0).astype(np.float64)

        ls_core = (weights_d.to_numpy(dtype=np.float64) * ret_filled.to_numpy(dtype=np.float64)).sum(axis=1)
        daily_rf = float(cash_annual_yield) / 252.0
        gross = weights_d.abs().sum(axis=1).astype(np.float64)
        cash_w = np.maximum(0.0, 2.0 - gross.to_numpy(dtype=np.float64))
        ls_daily_ret = ls_core + cash_w * daily_rf

        out = pd.Series(ls_daily_ret, index=master_index, name="ls_ret", dtype=np.float64)

        if verbose:
            print(
                f"  ML L/S: walk-forward month-ends={len(weights_m_df)}  "
                f"top/bottom n={self.top_n}  train_window≈{self.train_years}y  "
                f"target={tm}  weighting={wt}  features: SPY beta/corr/rel-mom + VIX"
            )

        return out


def build_ml_ls_returns(
    equity_dict: Dict[str, pd.DataFrame],
    *,
    sector_map: Dict[str, str] | None = None,
    cash_annual_yield: float = 0.04,
    engine: MLMomentumEngine | None = None,
    **engine_kwargs: Any,
) -> pd.Series:
    """
    Functional wrapper around :class:`MLMomentumEngine` for :class:`EnsembleManager` integration.

    Pass a custom ``engine`` or constructor fields such as ``train_years=3`` via ``engine_kwargs``.
    """
    allowed = ("train_years", "horizon", "top_n", "xgb_params", "data_period", "target_mode", "weighting")
    kw = {k: v for k, v in engine_kwargs.items() if k in allowed}
    eng = engine or MLMomentumEngine(**kw)
    return eng.generate_ls_returns(equity_dict, sector_map=sector_map, cash_annual_yield=cash_annual_yield)
