# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #
#                                                                             #
#   06-network-meta-analysis.R                                                #
#   AD network meta-analysis – continuous outcome                             #
#   No covariates, no moderators. Data: implist[[1]] (y complete).            #
#   Arm-level means ȳ[i,k] ~ N(mu[i]+delta[i,k], se²); Lu & Ades (2006) RE.   #
#   BUGS Book §11.3–11.4; BDA3 §5.6.                                          #
#   Output: mean differences d[t] (vs reference), forest plot.                #
#                                                                             #
# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #


# 0. Packages --------------------------------------------------------------

library(runjags)
library(coda)
library(dplyr)
library(ggplot2)
library(mitml)
library(purrr)
library(netmeta)
library(meta)
require(tidyr)

# 1. Global settings --------------------------------------------------------

path     = "_data/_imputation/implist.rda"
load(path)

set.seed(2026)
ref.trt  = 1L
n.chains = 3L
n.adapt  = 3000L
n.burn   = 10000L
n.iter   = 20000L
n.thin   = 10L
n.sets   = length(implist)

# 2. Compute arm-level summaries from IPD ------------------------------------
# y is assumed complete — all imputations give identical summaries.

# Step 1: per-imputation summaries
imp.stats = map(implist, function(x){
  x %>%
    filter(!is.na(y)) %>%
    group_by(study, treat) %>%
    summarise(
      mean.y = mean(y),
      var.y  = var(y),
      n      = n(),
      .groups = "drop"
    ) %>%
    arrange(study, treat)
})

# Step 2: pool via Rubin's Rules
bind_rows(imp.stats, .id = "imp") %>%
  group_by(study, treat) %>%
  summarise(
    n      = mean(n),
    B      = var(mean.y),              # between-imputation variance (divides by m-1)
    mean.y = mean(mean.y),             # Rubin: Q̄ = mean of Q̂_l
    W      = mean(var.y / n),          # within-imputation variance of mean
    T_var  = W + (1 + 1/n.sets) * B,   # Rubin's total variance
    se.y   = sqrt(T_var),              # SE of pooled mean
    sd.y   = sqrt(mean(var.y)),        # pooled SD of y (for reporting)
    .groups = "drop"
  ) %>%
  select(study, treat, mean.y, sd.y, se.y, T_var, n) %>%
  mutate(treat = treat %>% as.character %>% as.numeric) %>% 
  arrange(study, treat) -> arm.stats

cat("Arm-level summaries (first 12 rows):\n")
print(head(arm.stats, 12))


# 3. Study metadata ---------------------------------------------------------

study.info = arm.stats %>%
  group_by(study) %>%
  summarise(
    na       = n(),
    t.global = list(sort(unique(treat))),
    .groups  = "drop"
  ) %>%
  arrange(study)

ns         = nrow(study.info)
study.ids  = study.info$study
Nt         = max(arm.stats$treat)
arm.counts = sort(unique(study.info$na))

cat(sprintf("\nStudies: %d  |  Treatments: %d\n", ns, Nt))
cat("Studies per arm count:\n")
print(table(study.info$na))


# 4. Assemble JAGS data (grouped by arm count) ------------------------------

jags.data = list(Nt = Nt)

for (na.str in as.character(arm.counts)) {
  na  = as.integer(na.str)
  idx = which(study.info$na == na)
  N   = length(idx)

  y.mat    = matrix(NA_real_,    nrow = N, ncol = na)
  prec.mat = matrix(NA_real_,    nrow = N, ncol = na)
  t.mat    = matrix(NA_integer_, nrow = N, ncol = na)

  for (ii in seq_len(N)) {
    si = study.ids[idx[ii]]
    tg = study.info$t.global[[idx[ii]]]   # sorted global codes
    for (kk in seq_len(na)) {
      row = arm.stats %>% filter(study == si, treat == tg[kk])
      y.mat[ii, kk]    = row$mean.y
      prec.mat[ii, kk] = 1/(row$se.y^2)
    }
    t.mat[ii, ] = tg
  }

  jags.data[[paste0("N",    na.str)]] = N
  jags.data[[paste0("y.",   na.str)]] = y.mat
  jags.data[[paste0("prec.",na.str)]] = prec.mat
  jags.data[[paste0("t.",   na.str)]] = t.mat
}


# 5. JAGS model builder -----------------------------------------------------
# Generates a model string for any set of arm-count groups.
# Variables are named with arm-count suffix: y.2, prec.2, delta.2, mu.2, etc.

build.model = function(arm.counts) {
  L   = function(...) paste0(..., "\n")
  out = L("model {")
  out = paste0(out,
    L(""),
    L("  # Pooled mean differences (d[1] = 0 = reference treatment)"),
    L("  d[1] <- 0"),
    L("  for (t in 2:Nt) { d[t] ~ dnorm(0, 0.001) }"),
    L(""),
    L("  # Between-study heterogeneity in mean differences"),
    L("  tau      ~ dunif(0, 5)"),
    L("  prec.tau <- 1 / (tau * tau)"),
    L("")
  )

  for (na.str in as.character(sort(arm.counts))) {
    na   = as.integer(na.str)
    nm1  = na - 1
    Nv   = paste0("N",     na.str)
    yv   = paste0("y.",    na.str)
    pv   = paste0("prec.", na.str)
    tv   = paste0("t.",    na.str)
    muv  = paste0("mu.",   na.str)
    xiv  = paste0("xi.",   na.str)
    thv  = paste0("theta.",na.str)
    rcv  = paste0("prec.la.", na.str)

    out = paste0(out, L(sprintf("  ########## %d-arm studies ##########", na)))
    out = paste0(out, L(sprintf("  for (i in 1:%s) {", Nv)))

    # Study baseline (nuisance)
    out = paste0(out, L(sprintf("    %s[i] ~ dnorm(0, 0.0001)", muv)))

    # Arm-level likelihoods
    for (k in seq_len(na)) {
      if (k == 1) {
        out = paste0(out, L(sprintf("    %s[i,%d] ~ dnorm(%s[i],              %s[i,%d])", yv, k, muv, pv, k)))
      } else if (na == 2) {
        out = paste0(out, L(sprintf("    %s[i,%d] ~ dnorm(%s[i] + delta.2[i], %s[i,%d])", yv, k, muv, pv, k)))
      } else {
        out = paste0(out, L(sprintf("    %s[i,%d] ~ dnorm(%s[i] + %s[i,%d],   %s[i,%d])", yv, k, muv, xiv, k-1, pv, k)))
      }
    }

    # Random effects
    if (na == 2) {
      out = paste0(out,
        L(sprintf("    delta.2[i] ~ dnorm(d[%s[i,2]] - d[%s[i,1]], prec.tau)", tv, tv)))
    } else {
      out = paste0(out,
        L(sprintf("    %s[i,1:%d] ~ dmnorm(%s[i,1:%d], %s[1:%d,1:%d])", xiv, nm1, thv, nm1, rcv, nm1, nm1)),
        L(sprintf("    for (kk in 1:%d) {", nm1)),
        L(sprintf("      %s[i,kk] <- d[%s[i,kk+1]] - d[%s[i,1]]", thv, tv, tv)),
        L("    }")
      )
    }

    out = paste0(out, L("  }"), L(""))
  }

  # Lu-Ades precision matrices for na >= 3
  for (na.str in as.character(sort(arm.counts[arm.counts >= 3]))) {
    na   = as.integer(na.str)
    nm1  = na - 1
    rcv  = paste0("prec.la.", na.str)
    diag.v   =  2 * nm1 / na
    off.v    = -2 / na
    out = paste0(out, 
      L(sprintf("  # Lu-Ades precision (%d-arm): diag=%.4g*prec.tau, off=%.4g*prec.tau", 
      na, diag.v, off.v)))
    for (j in seq_len(nm1)) {
      for (l in seq_len(nm1)) {
        val = if (j == l) diag.v else off.v
        out = paste0(out, L(sprintf("  %s[%d,%d] <- %.10g * prec.tau", rcv, j, l, val)))
      }
    }
    out = paste0(out, L(""))
  }

  paste0(out, "}\n")
}

model.string = build.model(arm.counts)

# Print the create NMA model in JAGS code
cat("\n--- JAGS model ---\n")
cat(model.string)


# 6. Run JAGS ----------------------------------------------------------------

fit = run.jags(model = model.string, data = jags.data, 
               monitor = c("d", "tau"),
               n.chains = n.chains, adapt = n.adapt,
               burnin = n.burn, sample = n.iter, thin = n.thin, 
               summarise = FALSE)
summary(fit)
samp     = as.mcmc.list(fit)
post.mat = as.matrix(samp)

# Compare the results with a frequentist analysis
meta::pairwise(treat, n = n, mean = mean.y, sd = sd.y,
               studlab = study, data = arm.stats) %>% 
  netmeta::netmeta(TE, seTE, treat1, treat2, studlab, data = .)



# 7. Convergence diagnostics ------------------------------------------------

cat("\n--- Gelman-Rubin R-hat (target < 1.1) ---\n")
rhat = gelman.diag(samp, multivariate = FALSE)$psrf[, 1]
print(round(sort(rhat, decreasing = TRUE), 3))
cat(sprintf("%d / %d parameters with R-hat < 1.1\n",
            sum(rhat < 1.1, na.rm = TRUE), length(rhat)-1))

ess = effectiveSize(samp)
cat(sprintf("ESS — Min: %d  Median: %d\n", round(min(ess[-1])), round(median(ess[-1]))))


# 8. Posterior summary ------------------------------------------------------

posterior.summary = function(mat, params = NULL, prob = 0.95) {
  if (!is.null(params)) mat = mat[, intersect(params, colnames(mat)), drop = FALSE]
  as.data.frame(t(apply(mat, 2, function(x) {
    h = HDInterval::hdi(x, credMass = prob)
    c(
      mean   = mean(x),
      sd     = sd(x),
      hdi.lo = unname(h["lower"]),
      median = unname(quantile(x, 0.500)),
      hdi.hi = unname(h["upper"]),
      P.pos  = mean(x > 0)
    )
  })))
}

trt.labels = paste0("Trt", seq_len(Nt))
d.params   = paste0("d[", seq_len(Nt), "]")

cat("\n--- Pooled mean differences vs reference (d[t]) ---\n")
ps.d = posterior.summary(post.mat, d.params)
rownames(ps.d) = trt.labels
print(round(ps.d, 3))

cat("\n--- Between-study heterogeneity (tau) ---\n")
print(round(posterior.summary(post.mat, "tau"), 3))



# 9. Trace plots -----------------------------------------------------------

trace.params = c(paste0("d[", seq_len(min(4, Nt)), "]"), "tau")
trace.params = intersect(trace.params, colnames(post.mat))
par(mfrow = c(ceiling(length(trace.params) / 2), 2), mar = c(3, 3, 2, 1))
for (p in trace.params) traceplot(samp[, p], main = p, ask = FALSE)
par(mfrow = c(1, 1))


# 10. Treatment-effect league table -----------------------------------------
# All pairwise contrasts d[t] - d[s] from the posterior

league = expand.grid(t = seq_len(Nt), s = seq_len(Nt)) %>%
  filter(t != s) %>%
  rowwise() %>%
  mutate(
    contrast = sprintf("Trt%d vs Trt%d", t, s),
    draws    = list(post.mat[, paste0("d[", t, "]")] - post.mat[, paste0("d[", s, "]")])
  ) %>%
  ungroup() %>%
  mutate(
    mean   = sapply(draws, mean),
    hdi.lo = sapply(draws, function(x) HDInterval::hdi(x)["lower"]),
    hdi.hi = sapply(draws, function(x) HDInterval::hdi(x)["upper"]),
    P.pos  = sapply(draws, function(x) mean(x > 0))) %>%
  mutate(cell = sprintf("%.3f [%.3f, %.3f]", mean, hdi.lo, hdi.hi)) %>%
  select(t, s, cell) %>%
  tidyr::pivot_wider(names_from = s, values_from = cell) %>%
  rename_with(~paste0("vs Trt", .), -t) %>%
  mutate(t = paste0("Trt", t)) %>%
  arrange(t) %>% rename(` ` = t)  


