"""
Chan Johansen triplet mean-reversion (Ex 2.7–2.8).

Unit portfolio from first Johansen eigenvector on log prices; linear z-score
``numUnits = -(yport - MA) / STD`` with half-life lookback; dollar-position PnL.
"""

from __future__ import annotations

from typing import Any, Literal

import numpy as np
import pandas as pd

import universe_scanner as us  # type: ignore[import-not-found]

from RenTech.strategy_stack.statarb_engine import rolling_zscore

Fidelity = Literal["book", "causal"]

try:
    from statsmodels.tsa.vector_ar.vecm import coint_johansen
except ImportError:  # pragma: no cover
    coint_johansen = None  # type: ignore[misc, assignment]


def half_life_lookback(spread: pd.Series, *, default: int = 20, cap: int = 120) -> int:
    hl = us.calculate_half_life(spread)
    if np.isfinite(hl) and 1 <= hl <= cap:
        return max(5, int(round(hl)))
    return default


def chan_dollar_portfolio_returns(
    prices: pd.DataFrame,
    unit_weights: pd.DataFrame,
    num_units: pd.Series,
) -> pd.Series:
    """Ex 2.8 dollar-position return."""
    nu = num_units.astype(np.float64)
    w = unit_weights.reindex(prices.index).astype(np.float64)
    px = prices.astype(np.float64)
    dollar_pos = w.mul(px, axis=0).mul(nu, axis=0)
    pct = px.pct_change()
    pnl = (dollar_pos.shift(1) * pct).sum(axis=1)
    gross = dollar_pos.shift(1).abs().sum(axis=1).replace(0.0, np.nan)
    return (pnl / gross).fillna(0.0)


def linear_zscore_units(spread: pd.Series, lookback: int, *, fidelity: Fidelity) -> pd.Series:
    mu = spread.rolling(lookback, min_periods=max(2, lookback // 2)).mean()
    sig = spread.rolling(lookback, min_periods=max(2, lookback // 2)).std(ddof=1)
    z = (spread - mu) / sig.replace(0.0, np.nan)
    units = -z
    if fidelity == "causal":
        units = units.clip(-3.0, 3.0)
    return units.fillna(0.0)


def johansen_eigenvector_at(
    log_px: pd.DataFrame,
    t: int,
    *,
    fidelity: Fidelity,
    min_train: int,
    refit_bars: int,
) -> np.ndarray | None:
    if coint_johansen is None:
        raise ImportError("statsmodels required for Johansen tests")
    if fidelity == "book":
        window = log_px.iloc[: t + 1]
    else:
        if (t - min_train) % refit_bars != 0 and t != min_train:
            return None
        start = max(0, t - min_train + 1)
        window = log_px.iloc[start : t + 1]
    if len(window) < min_train:
        return None
    res = coint_johansen(window.to_numpy(), det_order=0, k_ar_diff=1)
    ev = res.evec[:, 0].astype(np.float64)
    if ev[0] != 0:
        ev = ev / ev[0]
    return ev


def johansen_trace_pvalue(log_px: pd.DataFrame) -> tuple[float, float, np.ndarray]:
    """
    Johansen trace test for r=0. Returns (trace_stat, crit_95, eigenvector).

    ``crit_95`` is the 95% critical value for rejecting r=0.
    """
    if coint_johansen is None:
        raise ImportError("statsmodels required")
    arr = log_px.dropna(how="any").astype(np.float64)
    if len(arr) < 60:
        raise ValueError("need >= 60 rows for Johansen")
    res = coint_johansen(np.log(arr).to_numpy(), det_order=0, k_ar_diff=1)
    trace = float(res.lr1[0])
    crit95 = float(res.cvt[0, 1])
    ev = res.evec[:, 0].astype(np.float64)
    if ev[0] != 0:
        ev = ev / ev[0]
    return trace, crit95, ev


def johansen_linear_triplet(
    prices: pd.DataFrame,
    *,
    fidelity: Fidelity = "causal",
    refit_bars: int = 63,
    min_train: int = 252,
) -> tuple[pd.Series, dict[str, Any]]:
    """Backtest one fixed triplet (columns = tickers)."""
    px = prices.dropna(how="any").astype(np.float64)
    cols = list(px.columns)
    log_px = np.log(px)
    n = len(px)
    evec = np.full((n, len(cols)), np.nan)
    last_ev: np.ndarray | None = None
    for i in range(min_train, n):
        ev = johansen_eigenvector_at(
            log_px, i, fidelity=fidelity, min_train=min_train, refit_bars=refit_bars
        )
        if ev is not None:
            last_ev = ev
        if last_ev is not None:
            evec[i] = last_ev
    evec_df = pd.DataFrame(evec, index=px.index, columns=cols).ffill()
    yport = (px * evec_df).sum(axis=1)
    lb = half_life_lookback(yport)
    units = linear_zscore_units(yport, lb, fidelity=fidelity)
    ret = chan_dollar_portfolio_returns(px, evec_df, units)
    return ret, {"lookback": lb, "min_train": min_train, "fidelity": fidelity, "legs": cols}


def unit_portfolio_half_life(prices: pd.DataFrame, evec: np.ndarray) -> float:
    px = prices.dropna(how="any").astype(np.float64)
    w = pd.Series(evec, index=px.columns)
    yport = (px * w).sum(axis=1)
    return float(us.calculate_half_life(yport))
