"""
Macro regime filter: time-series momentum on daily data.

Combines a long moving average (trend) with a medium-term rate of change (momentum).

Regime coding (default rules):
  +1  Bullish — price above SMA and ROC positive (both agree uptrend)
  -1  Bearish — price below SMA and ROC negative
   0  Flat / mixed — contradictory or not enough history
"""

from __future__ import annotations

from dataclasses import dataclass

import numpy as np
import pandas as pd


def compute_regime_series(
    close: pd.Series,
    *,
    sma_window: int = 200,
    aqr_lookback: int = 252,
    aqr_skip: int = 21,
) -> pd.Series:
    """
    Vectorized daily regime indicator.

    Parameters
    ----------
    close
        Daily close levels (DatetimeIndex).
    sma_window
        Simple moving average length (default 200 trading days).
    aqr_lookback
        AQR lookback in trading days (default 252 = ~12 months).
    aqr_skip
        AQR skip in trading days to exclude the most recent momentum window
        (default 21 = ~1 month).

    Returns
    -------
    Series of int8 in {-1, 0, 1}, index aligned to `close`.
    """
    c = close.astype(np.float64)
    sma = c.rolling(window=sma_window, min_periods=sma_window).mean()
    # AQR 12-minus-1 momentum:
    #   Return = Price_{t-skip} / Price_{t-lookback} - 1
    aqr_mom = c.shift(int(aqr_skip)) / c.shift(int(aqr_lookback)) - 1.0

    bullish = (c > sma) & (aqr_mom > 0.0)
    bearish = (c < sma) & (aqr_mom < 0.0)

    regime = pd.Series(np.int8(0), index=c.index)
    regime = regime.mask(bullish, np.int8(1))
    regime = regime.mask(bearish, np.int8(-1))
    # If both True (shouldn't happen with strict inequalities), prefer flat
    regime = regime.mask(bullish & bearish, np.int8(0))
    return regime.astype(np.int8)


@dataclass
class MomentumFilter:
    """
    Configurable wrapper around :func:`compute_regime_series`.

    Typical use with :class:`RenTech.strategy_stack.data_loader.DataLoader`:

        loader = DataLoader()
        spy = loader.fetch_daily("SPY", period="10y")
        filt = MomentumFilter()
        regime = filt.transform(spy["close"])
    """

    sma_window: int = 200
    aqr_lookback: int = 252
    aqr_skip: int = 21

    def transform(self, close: pd.Series) -> pd.Series:
        """Return regime Series {-1, 0, 1}."""
        return compute_regime_series(
            close,
            sma_window=self.sma_window,
            aqr_lookback=self.aqr_lookback,
            aqr_skip=self.aqr_skip,
        )

    def transform_frame(self, ohlcv: pd.DataFrame, *, price_col: str = "close") -> pd.DataFrame:
        """
        Append regime columns to an OHLCV DataFrame.

        Adds: `sma_{window}`, `aqr_mom_{aqr_lookback}_{aqr_skip}`, `regime`.
        """
        if price_col not in ohlcv.columns:
            raise KeyError(f"{price_col} not in frame columns: {list(ohlcv.columns)}")
        out = ohlcv.copy()
        c = out[price_col]
        out[f"sma_{self.sma_window}"] = c.rolling(self.sma_window, min_periods=self.sma_window).mean()
        out[f"aqr_mom_{self.aqr_lookback}_{self.aqr_skip}"] = (
            c.shift(int(self.aqr_skip)) / c.shift(int(self.aqr_lookback)) - 1.0
        )
        out["regime"] = self.transform(c)
        return out


if __name__ == "__main__":
    # Local smoke test without network: synthetic random walk
    rng = np.random.default_rng(42)
    idx = pd.date_range("2015-01-01", periods=300, freq="B")
    px = 100 * np.exp(np.cumsum(rng.normal(0, 0.008, size=len(idx))))
    s = pd.Series(px, index=idx)
    reg = compute_regime_series(s, sma_window=50, aqr_lookback=20, aqr_skip=5)
    print(reg.value_counts(dropna=False))
    print(reg.tail(10))
