"""Markov-switching VAR (Bayesian, from scratch).

A VAR whose intercept and shock covariance switch between K regimes governed by a latent Markov chain
(MSIH-VAR: switching intercept + heteroskedasticity, constant autoregressive dynamics).
Gibbs sampler:
  (1) regime path S_t   -- forward Hamilton filter + backward sampling (FFBS)
  (2) intercepts c_k, covariances Sigma_k, and the constant lag matrix B  -- conjugate Normal / inverse-Wishart / GLS
  (3) transition matrix P  -- Dirichlet
Self-contained (numpy only); alters nothing else.
"""
import numpy as np


def _lag(Y, p):
    T, m = Y.shape
    Yt = Y[p:]
    Xl = np.hstack([Y[p - l:T - l] for l in range(1, p + 1)])     # (T-p, m*p) lags, NO intercept
    return Yt, Xl


def _ffbs(dens, P, pi0, rng):
    """Forward Hamilton filter + backward sample.  dens (T,K)=p(y_t|S_t=k); returns sampled path S (T,)."""
    T, K = dens.shape
    filt = np.zeros((T, K)); pred = pi0.copy()
    for t in range(T):
        f = pred * dens[t]; f /= f.sum(); filt[t] = f
        pred = P.T @ f
    S = np.zeros(T, int)
    S[T - 1] = rng.choice(K, p=filt[T - 1])
    for t in range(T - 2, -1, -1):
        w = filt[t] * P[:, S[t + 1]]; w /= w.sum()
        S[t] = rng.choice(K, p=w)
    return S


def gibbs_msvar(Y, p, K=2, ndraw=3000, burn=2000, seed=1, iw_scale=1.0):
    """MSIH-VAR: switching intercept + covariance, constant autoregressive dynamics (stable & interpretable)."""
    rng = np.random.default_rng(seed)
    Yt, Xl = _lag(Y, p); n, m = Yt.shape; q = Xl.shape[1]
    B = np.linalg.lstsq(np.hstack([np.ones((n, 1)), Xl]), Yt, rcond=None)[0][1:]      # (q,m) constant lag matrix
    resid = Yt - Xl @ B
    S = (resid[:, 0] < np.median(resid[:, 0])).astype(int)
    c = np.array([resid[S == k].mean(0) for k in range(K)])                           # (K,m) intercepts
    Sig = np.array([np.cov(resid[S == k].T) for k in range(K)])
    P = np.full((K, K), 0.1) + np.eye(K) * 0.8
    nu0 = m + 6; S0 = iw_scale * np.diag(resid.var(0)) * nu0                          # data-scaled IW prior
    keepS, keepC, keepSig, keepP = [], [], [], []

    for it in range(ndraw + burn):
        # ---- (1) FFBS of the regime path ----
        dens = np.zeros((n, K))
        for k in range(K):
            e = Yt - c[k] - Xl @ B; Si = np.linalg.inv(Sig[k]); ld = np.linalg.slogdet(Sig[k])[1]
            dens[:, k] = np.exp(-0.5 * (np.einsum("ti,ij,tj->t", e, Si, e) + ld))
        dens = np.clip(dens, 1e-300, None)
        S = _ffbs(dens, P, np.ones(K) / K, rng)
        order = np.argsort([np.sqrt(Sig[k, 0, 0]) for k in range(K)])                 # regime 0 = calm (low vol), last = turbulent
        remap = np.argsort(order); S = remap[S]; c, Sig = c[order], Sig[order]; P = P[np.ix_(order, order)]

        # ---- (2) constant B via GLS across regimes ----
        XtX = np.zeros((q * m, q * m)); Xty = np.zeros(q * m)
        for t in range(n):
            Si = np.linalg.inv(Sig[S[t]])
            XtX += np.kron(np.outer(Xl[t], Xl[t]), Si); Xty += np.kron(Xl[t][:, None], Si) @ (Yt[t] - c[S[t]])
        XtX += np.eye(q * m) / 50.0
        Cov = np.linalg.inv(XtX); B = (Cov @ Xty + np.linalg.cholesky(Cov) @ rng.standard_normal(q * m)).reshape(q, m)
        # ---- intercepts & covariances per regime ----
        for k in range(K):
            e = (Yt - Xl @ B)[S == k]; nk = len(e)
            if nk > m + 1:
                Sk = np.linalg.inv(Sig[k]); pc = nk * Sk + np.eye(m) / 50.0
                c[k] = np.linalg.solve(pc, nk * Sk @ e.mean(0)) + np.linalg.cholesky(np.linalg.inv(pc)) @ rng.standard_normal(m)
                r = e - c[k]
                L = np.linalg.cholesky(np.linalg.inv(S0 + r.T @ r)); A = L @ rng.standard_normal((m, nu0 + nk))
                Sig[k] = np.linalg.inv(A @ A.T)
        # ---- (3) transition matrix ----
        for j in range(K):
            cnt = np.array([np.sum((S[:-1] == j) & (S[1:] == kk)) for kk in range(K)]) + 1.0
            g = rng.gamma(cnt); P[j] = g / g.sum()

        if it >= burn:
            keepS.append(S.copy()); keepC.append(c.copy()); keepSig.append(Sig.copy()); keepP.append(P.copy())
    return {"S": np.array(keepS), "c": np.array(keepC), "Sigma": np.array(keepSig),
            "P": np.array(keepP), "B": B, "p": p, "K": K, "Yt": Yt}
