"""Bayesian Vector Error Correction Model (VECM) — from scratch.

Reduced-rank cointegrated VAR in error-correction form
    dY_t = c + alpha beta' Y_{t-1} + sum_{i=1}^{p-1} Gamma_i dY_{t-i} + eps_t,
    eps_t ~ N(0, Sigma),  Pi = alpha beta' has rank r.

Identification: linear normalization beta = [I_r ; B], with B the free (m-r) x r block.
Estimation: 3-block Gibbs sampler (all blocks conjugate given the normalization):
    (1) Sigma | rest            ~ inverse-Wishart
    (2) [alpha, Theta] | beta   ~ matric-variate Normal   (Theta = [Gamma_1..Gamma_{p-1}, c])
    (3) B | alpha, Theta, Sigma ~ Normal                  (the EC term is linear in B)

This module is self-contained (numpy only) and does not touch bvar.py.
"""
import numpy as np


def vecm_matrices(Y, p):
    """Build VECM regression pieces from levels Y (T,m), p lags in levels (p-1 diff lags).
    Returns dYt (n,m)=Delta Y_t, Ylag (n,m)=Y_{t-1}, Xsr (n,g)=[Delta Y_{t-1..t-(p-1)}, const]."""
    T, m = Y.shape
    dY = np.diff(Y, axis=0)                         # dY[k] = Delta Y_{k+1}
    n = T - p
    dYt = dY[p - 1:]                                # Delta Y_t, t=p..T-1
    Ylag = Y[p - 1:T - 1]                           # Y_{t-1}
    parts = [dY[p - 1 - i:T - 1 - i] for i in range(1, p)]   # Delta Y_{t-i}
    Xsr = np.hstack(parts) if parts else np.empty((n, 0))
    Xsr = np.hstack([Xsr, np.ones((n, 1))])         # + intercept
    return dYt, Ylag, Xsr


def _beta_from_B(B, r, m):
    """beta = [I_r ; B]  (m x r), B is (m-r) x r."""
    beta = np.zeros((m, r))
    beta[:r] = np.eye(r)
    if m - r > 0:
        beta[r:] = B
    return beta


def johansen_init(dYt, Ylag, Xsr, r):
    """Reduced-rank (Johansen) OLS start: partial out short-run regressors, canonical correlations."""
    def resid(A, X):
        b, *_ = np.linalg.lstsq(X, A, rcond=None)
        return A - X @ b
    R0 = resid(dYt, Xsr)                            # Delta Y_t | short-run
    R1 = resid(Ylag, Xsr)                           # Y_{t-1}   | short-run
    n = R0.shape[0]
    S00 = R0.T @ R0 / n; S11 = R1.T @ R1 / n; S01 = R0.T @ R1 / n
    # eigen-problem S11^{-1} S10 S00^{-1} S01
    S11i = np.linalg.inv(S11)
    M = S11i @ S01.T @ np.linalg.inv(S00) @ S01
    w, V = np.linalg.eig(M)
    order = np.argsort(-w.real)
    V = V.real[:, order]
    beta = V[:, :r]
    # normalize to [I_r ; B]
    beta = beta @ np.linalg.inv(beta[:r])
    return beta, w.real[order]


def gibbs_vecm(Y, p, r, ndraw=4000, burn=2000, seed=1,
               b_prior_sd=10.0, c_prior_sd=10.0, iw_scale=1.0, B_prior_mean=None):
    """Gibbs sampler for the rank-r Bayesian VECM.  Returns dict of posterior draws."""
    rng = np.random.default_rng(seed)
    T, m = Y.shape
    dYt, Ylag, Xsr = vecm_matrices(Y, p)
    n, g = Xsr.shape
    mr = m - r                                      # free rows of beta

    beta, eig = johansen_init(dYt, Ylag, Xsr, r)
    B = beta[r:].copy() if mr > 0 else np.zeros((0, r))
    if B_prior_mean is None:
        B_prior_mean = np.zeros((mr, r))

    keepB, keepA, keepG, keepS = [], [], [], []
    Y1 = Ylag[:, :r]; Y2 = Ylag[:, r:]              # split levels for the EC term
    nu0 = m + 2; S0 = iw_scale * np.eye(m)

    def draw_iw(S, nu):
        Sinv = np.linalg.inv(S)
        L = np.linalg.cholesky(Sinv)
        A = L @ rng.standard_normal((m, nu))
        return np.linalg.inv(A @ A.T)

    for it in range(ndraw + burn):
        # ---- (2) [alpha, Theta] | beta, Sigma via GLS multivariate regression ----
        W = Ylag @ beta                             # (n,r) EC regressors
        R = np.hstack([W, Xsr])                     # (n, r+g)
        q = r + g
        # prior: vec(C) ~ N(0, diag(pr)); pr for alpha loose, Theta by c_prior_sd
        pr = np.concatenate([np.full(r, b_prior_sd**2), np.full(g, c_prior_sd**2)])
        if it == 0:
            Sig = np.cov((dYt - R @ np.linalg.lstsq(R, dYt, rcond=None)[0]).T)
        Sigi = np.linalg.inv(Sig)
        # posterior for C (m x q): vec(C') ~ N.  Precision = (R'R ⊗ Sigi) + diag(1/pr ⊗ Sigi?) -- use per-eq trick
        RtR = R.T @ R
        Prec = np.kron(RtR, Sigi) + np.kron(np.diag(1.0 / pr), np.eye(m))
        rhs = (Sigi @ dYt.T @ R).flatten(order="F")     # vec(Sigma^{-1} dY' R), column-major
        Cov = np.linalg.inv(Prec)
        cvec = Cov @ rhs + np.linalg.cholesky(Cov) @ rng.standard_normal(m * q)
        C = cvec.reshape(m, q, order="F")
        alpha = C[:, :r]; Theta = C[:, r:]

        # ---- (1) Sigma | rest ----
        E = dYt - R @ C.T
        Sig = draw_iw(S0 + E.T @ E, nu0 + n)

        # ---- (3) B | alpha, Theta, Sigma  (EC term linear in B) ----
        if mr > 0:
            eps_star = dYt - Xsr @ Theta.T - Y1 @ alpha.T      # remove short-run + fixed EC part
            Sigi = np.linalg.inv(Sig)
            # eps_star_t = (Y2_t' ⊗ alpha) vec(B') + eps ; stack GLS
            db = r * mr
            Prec_b = np.eye(db) / (b_prior_sd**2)
            rhs_b = (B_prior_mean.T.flatten(order="F")) / (b_prior_sd**2)
            for t in range(n):
                Mt = np.kron(Y2[t][None, :], alpha)            # m x (r*mr)  for vec(B'), B' is r x mr
                Prec_b += Mt.T @ Sigi @ Mt
                rhs_b += Mt.T @ Sigi @ eps_star[t]
            Covb = np.linalg.inv(Prec_b)
            bvec = Covb @ rhs_b + np.linalg.cholesky(Covb) @ rng.standard_normal(db)
            Bt = bvec.reshape(r, mr, order="F")                 # this is B' (column-major vec)
            B = Bt.T
            beta = _beta_from_B(B, r, m)

        if it >= burn:
            keepB.append(beta.copy()); keepA.append(alpha.copy())
            keepG.append(Theta.copy()); keepS.append(Sig.copy())

    return {"beta": np.array(keepB), "alpha": np.array(keepA), "Theta": np.array(keepG),
            "Sigma": np.array(keepS), "p": p, "r": r, "eig": eig, "m": m,
            "dYt": dYt, "Ylag": Ylag, "Xsr": Xsr}


def ec_term(fit, Y):
    """Posterior median error-correction series beta' Y_{t-1}  (n, r)."""
    beta = np.median(fit["beta"], axis=0)
    return fit["Ylag"] @ beta


def forecast_vecm(fit, Y, H, seed=1, npath=None):
    """Predictive level paths (ndraw|npath, H, m) by iterating the implied level dynamics."""
    rng = np.random.default_rng(seed)
    p, r, m = fit["p"], fit["r"], fit["m"]
    nd = fit["beta"].shape[0]
    idx = range(nd) if npath is None else rng.integers(0, nd, npath)
    out = []
    for s in idx:
        beta, alpha, Theta, Sig = fit["beta"][s], fit["alpha"][s], fit["Theta"][s], fit["Sigma"][s]
        Gam = [Theta[:, (i * m):(i * m + m)] for i in range(p - 1)]
        c = Theta[:, -1]
        L = np.linalg.cholesky(Sig)
        ylev = list(Y[-p:])                          # last p levels
        dhist = [ylev[-i] - ylev[-i - 1] for i in range(1, p)]   # recent diffs (newest first)
        path = []
        for h in range(H):
            dy = c + alpha @ (beta.T @ ylev[-1])
            for i in range(p - 1):
                dy = dy + Gam[i] @ dhist[i]
            dy = dy + L @ rng.standard_normal(m)
            ynew = ylev[-1] + dy
            path.append(ynew); ylev.append(ynew)
            dhist = [dy] + dhist[:-1] if p > 1 else []
        out.append(path)
    return np.array(out)
