"""Hierarchical natural-conjugate BVAR (Giannone-Lenza-Primiceri 2015) — from scratch.

The Minnesota prior is implemented by DUMMY OBSERVATIONS (Theil mixed estimation), which makes
the prior natural-conjugate (Normal-inverse-Wishart) and the marginal likelihood p(Y|hyper)
available in CLOSED FORM.  Hyperparameters:
    lam   - overall Minnesota tightness (smaller = tighter)
    mu    - sum-of-coefficients ("no-cointegration") prior loosening
    delta - dummy-initial-observation ("single-unit-root") prior loosening
Estimating these by maximising / sampling the marginal likelihood is the GLP hierarchical prior.
Self-contained (numpy/scipy only); does not touch bvar.py.
"""
import numpy as np
from scipy.special import multigammaln, gammaln


def ar_sigma(Y, p):
    """Scale s_i = residual std of an AR(p) fit of each series on its own lags (+const)."""
    T, m = Y.shape
    s = np.empty(m)
    for i in range(m):
        yi = Y[p:, i]
        Xi = np.column_stack([np.ones(T - p)] + [Y[p - l:T - l, i] for l in range(1, p + 1)])
        b, *_ = np.linalg.lstsq(Xi, yi, rcond=None)
        s[i] = (yi - Xi @ b).std(ddof=1)
    return s


def _lagmat(Y, p):
    """Actual regression data with intercept LAST: X=[y_{t-1},...,y_{t-p},1], Yt=y_t."""
    T, m = Y.shape
    Yt = Y[p:]
    X = np.hstack([Y[p - l:T - l] for l in range(1, p + 1)] + [np.ones((T - p, 1))])
    return Yt, X


def dummies(Y, p, lam, mu, delta, s, b1=1.0, eps=1e-4):
    """Minnesota + sum-of-coefficients + dummy-initial-observation dummy observations.
    Returns (Yd, Xd) with the same column layout as _lagmat (intercept last)."""
    T, m = Y.shape
    k = m * p + 1
    ybar = Y[:p].mean(0)                                             # pre-sample mean (m,)
    J = np.arange(1, p + 1)                                          # lag-decay factors (lam3=1)

    # --- Minnesota: prior on AR coefficients (m*p rows) ---
    Xmn = np.zeros((m * p, k)); Ymn = np.zeros((m * p, m))
    for l in range(p):
        Xmn[l * m:(l + 1) * m, l * m:(l + 1) * m] = np.diag(s) * J[l] / lam
    Ymn[:m] = np.diag(b1 * s) / lam                                 # prior mean b1 (=1: random walk) on own 1st lag
    # --- prior on residual covariance (m rows) ---
    Xcov = np.zeros((m, k)); Ycov = np.diag(s)
    # --- loose prior on the intercept (1 row) ---
    Xint = np.zeros((1, k)); Xint[0, -1] = eps; Yint = np.zeros((1, m))
    # --- sum-of-coefficients prior (m rows), hyperparameter mu ---
    Xsoc = np.zeros((m, k)); Ysoc = np.diag(ybar) * mu
    for l in range(p):
        Xsoc[:, l * m:(l + 1) * m] = np.diag(ybar) * mu
    # --- dummy-initial-observation / single-unit-root (1 row), hyperparameter delta ---
    Xdio = np.zeros((1, k)); Ydio = (ybar * delta)[None, :]
    for l in range(p):
        Xdio[0, l * m:(l + 1) * m] = ybar * delta
    Xdio[0, -1] = delta

    Yd = np.vstack([Ymn, Ycov, Yint, Ysoc, Ydio])
    Xd = np.vstack([Xmn, Xcov, Xint, Xsoc, Xdio])
    return Yd, Xd


def fit(Y, p, lam, mu, delta, s=None, b1=1.0):
    """Natural-conjugate posterior + log marginal likelihood via dummy observations."""
    if s is None:
        s = ar_sigma(Y, p)
    Yt, X = _lagmat(Y, p)
    Yd, Xd = dummies(Y, p, lam, mu, delta, s, b1)
    Ys = np.vstack([Yd, Yt]); Xs = np.vstack([Xd, X])              # augmented (prior + data)
    T = Yt.shape[0]; m = Y.shape[1]; k = X.shape[1]
    Td = Yd.shape[0]; d = Td - k                                    # prior inverse-Wishart dof

    XdXd = Xd.T @ Xd; XsXs = Xs.T @ Xs
    XsXs_i = np.linalg.inv(XsXs)
    Bhat = XsXs_i @ (Xs.T @ Ys)                                     # posterior mean coefficients
    Ed = Yd - Xd @ (np.linalg.solve(XdXd, Xd.T @ Yd))
    Sd = Ed.T @ Ed                                                  # prior scale
    Es = Ys - Xs @ Bhat
    Ss = Es.T @ Es                                                  # posterior scale

    # --- closed-form log marginal likelihood (Giannone-Lenza-Primiceri 2015) ---
    (s1, ld_XdXd) = np.linalg.slogdet(XdXd)
    (s2, ld_XsXs) = np.linalg.slogdet(XsXs)
    (s3, ld_Sd) = np.linalg.slogdet(Sd)
    (s4, ld_Ss) = np.linalg.slogdet(Ss)
    log_ml = (-m * T / 2.0 * np.log(np.pi)
              + multigammaln((T + d) / 2.0, m) - multigammaln(d / 2.0, m)
              - m / 2.0 * (ld_XsXs - ld_XdXd)
              + d / 2.0 * ld_Sd - (T + d) / 2.0 * ld_Ss)
    return {"Bhat": Bhat, "Ss": Ss, "XsXs_i": XsXs_i, "nu": T + d, "T": T, "m": m, "k": k,
            "p": p, "s": s, "log_ml": float(log_ml)}


def log_ml(Y, p, lam, mu, delta, s=None, b1=1.0):
    return fit(Y, p, lam, mu, delta, s, b1)["log_ml"]


def _mvt_logpdf(y, mu, Sig, nu):
    """log multivariate-t density, dim m, dof nu, location mu, scale matrix Sig."""
    m = len(y); d = y - mu
    sgn, ld = np.linalg.slogdet(Sig)
    q = d @ np.linalg.solve(Sig, d)
    return (gammaln((nu + m) / 2.0) - gammaln(nu / 2.0) - 0.5 * m * np.log(nu * np.pi)
            - 0.5 * ld - (nu + m) / 2.0 * np.log1p(q / nu))


def predict1(fit_, xnew):
    """One-step predictive of the conjugate BVAR: multivariate-t (mean, scale, dof)."""
    m, nu = fit_["m"], fit_["nu"]
    mean = xnew @ fit_["Bhat"]
    fac = 1.0 + xnew @ fit_["XsXs_i"] @ xnew
    nu_t = nu - m + 1
    Sig = fac * fit_["Ss"] / nu_t
    return mean, Sig, nu_t


def predictive_logscore(fit_, xnew, ynew):
    mean, Sig, nu_t = predict1(fit_, xnew)
    return _mvt_logpdf(ynew, mean, Sig, nu_t)


def _gamma_logpdf_modesd(x, mode, sd):
    """log Gamma(x) reparameterised by (mode, sd)  [GLP-style hyperpriors]."""
    from scipy.stats import gamma
    th = (np.sqrt(mode ** 2 + 4 * sd ** 2) - mode) / 2.0        # scale
    ksh = mode / th + 1.0                                        # shape
    return gamma.logpdf(x, ksh, scale=th)


def log_post_hyper(Y, p, lam, mu, delta, s):
    """log marginal likelihood + GLP Gamma hyperpriors (lambda mode .2 sd .4; mu,delta mode 1 sd 1)."""
    lp = (_gamma_logpdf_modesd(lam, 0.2, 0.4) + _gamma_logpdf_modesd(mu, 1.0, 1.0)
          + _gamma_logpdf_modesd(delta, 1.0, 1.0))
    if not np.isfinite(lp):
        return -np.inf
    return log_ml(Y, p, lam, mu, delta, s) + lp


def sample_hyper(Y, p, ndraw=2000, burn=1000, seed=1, step=0.10, s=None):
    """Random-walk Metropolis over log(lambda, mu, delta): the GLP hierarchical prior."""
    rng = np.random.default_rng(seed)
    if s is None:
        s = ar_sigma(Y, p)
    th = np.log(np.array([0.2, 1.0, 1.0]))
    lpost = log_post_hyper(Y, p, *np.exp(th), s)
    keep = []
    for it in range(ndraw + burn):
        prop = th + step * rng.standard_normal(3)
        lp = log_post_hyper(Y, p, *np.exp(prop), s)
        if np.log(rng.random()) < lp - lpost:
            th, lpost = prop, lp
        if it >= burn:
            keep.append(np.exp(th))
    return np.array(keep)                                        # (ndraw, 3): columns lambda, mu, delta


def _draw_iw(S, nu, rng):
    """Sigma ~ inverse-Wishart(scale S, dof nu)."""
    m = S.shape[0]
    L = np.linalg.cholesky(np.linalg.inv(S))
    A = L @ rng.standard_normal((m, int(nu)))
    return np.linalg.inv(A @ A.T)


def forecast_paths(Y, p, hyper, H, npath_per=20, seed=1):
    """Simulate H-step predictive LEVEL paths, integrating hyperparameter + coefficient + shock uncertainty.
    hyper is an array of (lambda, mu, delta) draws.  Returns (n_paths, H, m)."""
    rng = np.random.default_rng(seed)
    paths = []
    for (lam, mu, dl) in hyper:
        f = fit(Y, p, lam, mu, dl)
        Bhat, Ss, XsXs_i, nu, m = f["Bhat"], f["Ss"], f["XsXs_i"], f["nu"], f["m"]
        Lu = np.linalg.cholesky(XsXs_i)                          # k x k  (row covariance factor)
        for _ in range(npath_per):
            Sig = _draw_iw(Ss, nu, rng)
            Lv = np.linalg.cholesky(Sig)                         # m x m
            B = Bhat + Lu @ rng.standard_normal(Bhat.shape) @ Lv.T   # matric-normal draw of coefficients
            hist = [Y[-p + i] for i in range(p)]                 # last p levels, oldest..newest
            path = []
            for h in range(H):
                x = np.hstack([hist[-1 - l] for l in range(p)] + [1.0])
                y = x @ B + Lv @ rng.standard_normal(m)
                path.append(y); hist.append(y)
            paths.append(path)
    return np.array(paths)                                       # (n_paths, H, m)

