"""
Mean-reversion engines:

* **Single-asset** — rolling z-score on one ``close`` series (:class:`StatArbEngine`).
* **Pairs / stat-arb** — Engle–Granger-style hedge ratio (OLS), ADF on the residual
  spread, then rolling z-score on that spread (:class:`PairStatArbEngine`).

Dependencies: pandas, numpy; pairs mode also requires **statsmodels**.

:class:`PairStatArbEngine` fits the hedge **only on past data** (expanding or
rolling window), forms a **causal** spread at each bar, runs ADF on that spread,
and sets **basket PnL** to ``log_ret_y - β_{t-1}·log_ret_x`` (or simple returns
when not using logs). A full-sample OLS is still stored on the instance for
diagnostics (``ols_result``) but does not drive signals.
"""

from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any

import numpy as np
import pandas as pd

try:
    import statsmodels.api as sm
    from statsmodels.tsa.stattools import adfuller
except ImportError as e:  # pragma: no cover
    sm = None  # type: ignore[assignment]
    adfuller = None  # type: ignore[assignment]
    _SM_ERROR = e
else:
    _SM_ERROR = None


def _require_statsmodels() -> None:
    if sm is None or adfuller is None:
        raise ImportError(
            "statsmodels is required for pair / ADF logic. "
            "Install with: python -m pip install statsmodels\n"
            f"Original error: {_SM_ERROR}"
        )


def rolling_zscore(series: pd.Series, window: int, *, min_periods: int | None = None) -> pd.Series:
    """
    Z-score of the current value vs trailing window mean and std dev.

    Uses sample std (ddof=1) inside the window.
    """
    mp = min_periods if min_periods is not None else max(2, window // 2)
    x = series.astype(np.float64)
    mu = x.rolling(window=window, min_periods=mp).mean()
    sig = x.rolling(window=window, min_periods=mp).std(ddof=1)
    z = (x - mu) / sig.replace(0.0, np.nan)
    return z


def run_mean_reversion_positions(
    z: pd.Series,
    *,
    entry_z: float = 2.0,
    exit_z: float = 0.0,
) -> pd.Series:
    """
    Bar-by-bar position state machine on z-scores.

    Flat -> long (+1) when z <= -entry_z; flat -> short (-1) when z >= entry_z.
    Long exits when z >= exit_z; short exits when z <= exit_z.
    """
    pos = np.zeros(len(z), dtype=np.int8)
    current = np.int8(0)
    idx = z.index
    vals = z.to_numpy()

    for i in range(len(vals)):
        zi = vals[i]
        if not np.isfinite(zi):
            pos[i] = current
            continue

        if current == 0:
            if zi <= -entry_z:
                current = np.int8(1)
            elif zi >= entry_z:
                current = np.int8(-1)
        elif current == 1:
            if zi >= exit_z:
                current = np.int8(0)
        else:
            if zi <= exit_z:
                current = np.int8(0)

        pos[i] = current

    return pd.Series(pos, index=idx, dtype=np.int8)


def merge_intraday_pair(
    ohlcv_y: pd.DataFrame,
    ohlcv_x: pd.DataFrame,
    *,
    price_col: str = "close",
) -> pd.DataFrame:
    """Inner-join two OHLCV frames on timestamp; ``close_y`` / ``close_x``."""
    y = ohlcv_y[[price_col]].rename(columns={price_col: "close_y"})
    x = ohlcv_x[[price_col]].rename(columns={price_col: "close_x"})
    merged = y.join(x, how="inner")
    if merged.empty:
        raise ValueError("No overlapping timestamps between the two intraday series.")
    return merged


def causal_ols_spread(
    close_y: pd.Series,
    close_x: pd.Series,
    *,
    use_log: bool = True,
    hedge_window: int | None = None,
    min_train: int = 30,
) -> tuple[pd.Series, pd.Series, pd.Series, pd.Series]:
    """
    Causal OLS hedge: at bar *i*, fit ``Y ~ 1 + X`` using only **prior** rows
    ``[start, i-1]`` (expanding if ``hedge_window`` is ``None``, else rolling of
    length ``hedge_window``). Spread at *i* is the one-step-ahead residual
    ``y_i - (α + β x_i)`` with ``(α, β)`` from that fit.

    Returns
    -------
    spread, fitted, alpha, beta
        All ``pd.Series`` aligned to the common index (NaN until ``min_train``
        bars of history exist in the training slice).
    """
    _require_statsmodels()
    if hedge_window is not None and hedge_window <= 0:
        hedge_window = None
    if hedge_window is not None:
        min_train = min(int(min_train), int(hedge_window))
    df = pd.DataFrame({"close_y": close_y, "close_x": close_x}).dropna()
    if len(df) < min_train + 1:
        raise ValueError(f"Need at least min_train+1 overlapping bars; got {len(df)}")

    y_raw = np.log(df["close_y"]) if use_log else df["close_y"].astype(np.float64)
    x_raw = np.log(df["close_x"]) if use_log else df["close_x"].astype(np.float64)
    yv = y_raw.to_numpy()
    xv = x_raw.to_numpy()
    n = len(df)
    spread = np.full(n, np.nan)
    fitted = np.full(n, np.nan)
    alpha_a = np.full(n, np.nan)
    beta_a = np.full(n, np.nan)

    for i in range(n):
        if hedge_window is not None:
            start = max(0, i - int(hedge_window))
        else:
            start = 0
        end = i
        if end - start < min_train:
            continue
        y_train = yv[start:end]
        x_train = xv[start:end]
        X_train = np.column_stack((np.ones(end - start), x_train))
        try:
            coef, _, _, _ = np.linalg.lstsq(X_train, y_train, rcond=None)
        except Exception:
            continue
        a = float(coef[0])
        b = float(coef[1])
        yi, xi = float(yv[i]), float(xv[i])
        alpha_a[i] = a
        beta_a[i] = b
        fitted[i] = a + b * xi
        spread[i] = yi - fitted[i]

    idx = df.index
    return (
        pd.Series(spread, index=idx),
        pd.Series(fitted, index=idx),
        pd.Series(alpha_a, index=idx),
        pd.Series(beta_a, index=idx),
    )


def engle_granger_spread(
    close_y: pd.Series,
    close_x: pd.Series,
    *,
    use_log: bool = True,
) -> tuple[pd.Series, pd.Series, float, float, Any]:
    """
    First-stage Engle–Granger regression: regress Y on X (with constant).

    With logs: log(Y) = alpha + beta log(X) + epsilon; spread = epsilon (OLS residual).

    Returns
    -------
    spread
        Cointegrating residual series (aligned to common index).
    fitted
        In-sample fitted values (same index).
    alpha, beta
        Intercept and hedge coefficient on X.
    ols_result
        statsmodels ``OLS`` results object.
    """
    _require_statsmodels()
    df = pd.DataFrame({"close_y": close_y, "close_x": close_x}).dropna()
    if len(df) < 30:
        raise ValueError(f"Need more overlapping bars for regression; got {len(df)}")

    y = np.log(df["close_y"]) if use_log else df["close_y"].astype(np.float64)
    x = np.log(df["close_x"]) if use_log else df["close_x"].astype(np.float64)
    X = sm.add_constant(x)
    ols = sm.OLS(y, X).fit()
    fitted = pd.Series(ols.fittedvalues, index=df.index)
    spread = pd.Series(ols.resid, index=df.index)
    alpha = float(ols.params.iloc[0])
    beta = float(ols.params.iloc[1])
    return spread, fitted, alpha, beta, ols


def adf_on_spread(
    spread: pd.Series,
    *,
    autolag: str | None = "AIC",
    regression: str = "c",
    maxlag: int | None = None,
) -> dict[str, Any]:
    """
    Augmented Dickey–Fuller test on the spread (stationarity of residual).

    ``regression`` is passed to ``adfuller`` (default ``'c'`` = constant only).
    """
    _require_statsmodels()
    s = spread.dropna().astype(np.float64)
    if len(s) < 30:
        return {"error": "too_few_obs", "n": int(len(s))}

    kw: dict[str, Any] = {"regression": regression}
    if autolag is not None:
        kw["autolag"] = autolag
    if maxlag is not None:
        kw["maxlag"] = maxlag

    stat, pvalue, usedlag, nobs, crit, icbest = adfuller(s, **kw)
    crit = dict(crit) if crit is not None else {}
    return {
        "adf_statistic": float(stat),
        "pvalue": float(pvalue),
        "used_lag": int(usedlag),
        "n_obs": int(nobs),
        "critical_values": crit,
        "is_stationary_5pct": bool(stat < crit.get("5%", float("nan"))),
        "ic_best": icbest,
    }


@dataclass
class StatArbEngine:
    """
    Single-asset rolling z-score mean reversion on ``close``.
    """

    window: int = 20
    entry_z: float = 2.0
    exit_z: float = 0.0
    price_col: str = "close"

    def transform(self, ohlcv: pd.DataFrame) -> pd.DataFrame:
        if self.price_col not in ohlcv.columns:
            raise KeyError(f"{self.price_col} not in columns: {list(ohlcv.columns)}")
        out = ohlcv.copy()
        z = rolling_zscore(out[self.price_col], self.window)
        out["zscore"] = z
        out["micro_position"] = run_mean_reversion_positions(
            z, entry_z=self.entry_z, exit_z=self.exit_z
        )
        return out


@dataclass
class PairStatArbEngine:
    """
    Two-leg stat-arb: **causal** OLS hedge (expanding or rolling), ADF on that
    spread, rolling z-score, and **basket** bar returns for PnL.

    The **dependent** asset (Y) should match the name you use for downstream PnL
    (e.g. ``trade_ticker``); **X** is the hedge leg (e.g. sector ETF or basket).

    Columns added to the returned frame (aligned to Y after inner join):
      ``close_x``, ``pair_spread``, ``spread_fitted``, ``hedge_alpha``, ``hedge_beta``
      (time-varying), ``zscore``, ``micro_position``, ``spread_ret`` (``diff`` of
      causal spread), component returns, and ``basket_ret`` — the hedgeable
      return ``r_y - β_{t-1} r_x`` (log or simple, matching ``use_log_prices``).
    """

    window: int = 20
    entry_z: float = 2.0
    exit_z: float = 0.0
    use_log_prices: bool = True
    hedge_window: int | None = None
    min_hedge_obs: int = 30
    adf_autolag: str | None = "AIC"
    adf_regression: str = "c"
    price_col: str = "close"

    adf_report: dict[str, Any] = field(default_factory=dict, init=False, repr=False)
    ols_result: Any = field(default=None, init=False, repr=False)

    def transform(
        self,
        ohlcv_y: pd.DataFrame,
        ohlcv_x: pd.DataFrame,
        *,
        attach_y_ohlcv: bool = True,
    ) -> pd.DataFrame:
        """
        Causal spread from past-only OLS; ADF on that spread; z-score; MR positions.

        ``ols_result`` is set to **full-sample** Engle–Granger OLS for diagnostics
        only. Signals and ``adf_report`` use the **causal** spread series.

        If ``attach_y_ohlcv``, start from ``ohlcv_y`` rows that align with the merge;
        otherwise return only merge columns + signals.
        """
        py = self.price_col
        if py not in ohlcv_y.columns or py not in ohlcv_x.columns:
            raise KeyError(f"{py} required on both frames")

        merged = merge_intraday_pair(ohlcv_y, ohlcv_x, price_col=py)
        *_, ols_fs = engle_granger_spread(
            merged["close_y"],
            merged["close_x"],
            use_log=self.use_log_prices,
        )
        self.ols_result = ols_fs

        hw = self.hedge_window
        if hw is not None and hw <= 0:
            hw = None
        spread, fitted, alpha_s, beta_s = causal_ols_spread(
            merged["close_y"],
            merged["close_x"],
            use_log=self.use_log_prices,
            hedge_window=hw,
            min_train=self.min_hedge_obs,
        )
        self.adf_report = adf_on_spread(spread, autolag=self.adf_autolag, regression=self.adf_regression)

        z = rolling_zscore(spread, self.window)
        pos = run_mean_reversion_positions(z, entry_z=self.entry_z, exit_z=self.exit_z)

        if attach_y_ohlcv:
            out = ohlcv_y.reindex(spread.index).copy()
        else:
            out = pd.DataFrame(index=spread.index)

        cy = merged["close_y"].reindex(out.index)
        cx = merged["close_x"].reindex(out.index)
        hb = beta_s.reindex(out.index)

        out["close_x"] = cx
        out["pair_spread"] = spread.reindex(out.index)
        out["spread_fitted"] = fitted.reindex(out.index)
        out["hedge_alpha"] = alpha_s.reindex(out.index)
        out["hedge_beta"] = hb
        out["zscore"] = z.reindex(out.index)
        out["micro_position"] = pos.reindex(out.index).fillna(0).astype(np.int8)
        out["spread_ret"] = out["pair_spread"].diff()

        if self.use_log_prices:
            out["ret_y"] = np.log(cy / cy.shift(1))
            out["ret_x"] = np.log(cx / cx.shift(1))
        else:
            out["ret_y"] = cy.pct_change()
            out["ret_x"] = cx.pct_change()
        out["basket_ret"] = out["ret_y"] - hb.shift(1) * out["ret_x"]

        return out


if __name__ == "__main__":
    rng = np.random.default_rng(0)
    idx = pd.date_range("2024-01-01", periods=500, freq="h")
    # Cointegrated-ish random walk pair
    x = 100 + np.cumsum(rng.normal(0, 0.2, size=len(idx)))
    y = 0.5 * x + 10 + np.cumsum(rng.normal(0, 0.05, size=len(idx)))
    dfy = pd.DataFrame({"close": y}, index=idx)
    dfx = pd.DataFrame({"close": x}, index=idx)
    eng = PairStatArbEngine(window=40, entry_z=1.5, exit_z=0.0, min_hedge_obs=40)
    out = eng.transform(dfy, dfx)
    print("ADF (causal spread):", eng.adf_report)
    print(out[["close", "close_x", "pair_spread", "basket_ret", "zscore", "micro_position"]].dropna().head(8))
