# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #
#                                                                             #
#   06-network-meta-analysis-binary.R                                         #
#   AD network meta-analysis – binary outcome                                 #
#   Binary y created via within-study overall median split.                   #
#   Arm-level: r[i,k] / n[i,k]; Binomial likelihood, logit link.              #
#   Lu & Ades (2006) RE on log-OR scale. BDA3 §5.6; BUGS Book §11.3–11.4.     #
#   Output: log-ORs d[t] (vs reference), OR league table.                     #
#                                                                             #
# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #


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

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


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

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

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


# 2. Create binary y and compute arm-level (r, n) ------------------------------
# Within each study, compute the overall median of y pooled across all arms.
# y.bin = 1 if y > median(y in study), 0 otherwise.
# Since y is complete, binary y is identical across imputations; we average
# r and n across imputations for robustness and consistency with the pipeline.

imp.stats = map(implist, function(x) {
  x %>%
    filter(!is.na(y)) %>%
    group_by(study) %>%
    mutate(y.bin = as.integer(y > median(y))) %>%
    ungroup() %>%
    group_by(study, treat) %>%
    summarise(
      r = sum(y.bin),
      n = n(),
      .groups = "drop"
    ) %>%
    arrange(study, treat)
})

# Pool across imputations (averaging r and n; identical if y is complete)
bind_rows(imp.stats, .id = "imp") %>%
  group_by(study, treat) %>%
  summarise(
    r = round(mean(r)),
    n = round(mean(n)),
    .groups = "drop"
  ) %>%
  mutate(
    treat    = as.character(treat) %>% as.numeric,
    prop     = r / n,
    # Haldane-Anscombe correction (+0.5) guards against r = 0 or r = n
    log.or   = log((r + 0.5) / (n - r + 0.5)),
    se.log.or = sqrt(1 / (r + 0.5) + 1 / (n - r + 0.5))
  ) %>%
  arrange(study, treat) -> arm.stats

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

# Warn about sparse cells
sparse = arm.stats %>% filter(r == 0 | r == n)
if (nrow(sparse) > 0) {
  cat(sprintf("\nWARNING: %d arm(s) with zero or complete events (Haldane correction applied):\n", nrow(sparse)))
  print(sparse)
}


# 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 --------------------------------------------------------
# Replace continuous y/prec matrices with integer r/n matrices.

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)

  r.mat = matrix(NA_integer_, nrow = N, ncol = na)
  n.mat = matrix(NA_integer_, 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]]]
    for (kk in seq_len(na)) {
      row           = arm.stats %>% filter(study == si, treat == tg[kk])
      r.mat[ii, kk] = row$r
      n.mat[ii, kk] = row$n
    }
    t.mat[ii, ] = tg
  }

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


# 5. JAGS model builder (binomial likelihood, logit link) ----------------------

build.model = function(arm.counts) {
  L   = function(...) paste0(..., "\n")
  out = L("model {")
  out = paste0(out,
    L(""),
    L("  # Pooled log-ORs vs reference treatment (d[1] = 0)"),
    L("  d[1] <- 0"),
    L("  for (t in 2:Nt) { d[t] ~ dnorm(0, 0.001) }"),
    L(""),
    L("  # Between-study heterogeneity in log-ORs"),
    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)
    rv  = paste0("r.",     na.str)
    nv  = paste0("n.",     na.str)
    tv  = paste0("t.",     na.str)
    muv = paste0("mu.",    na.str)
    pv  = paste0("p.",     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 log-odds (nuisance)
    out = paste0(out, L(sprintf("    %s[i] ~ dnorm(0, 0.0001)", muv)))

    # Binomial likelihoods
    for (k in seq_len(na)) {
      out = paste0(out, L(sprintf(
        "    %s[i,%d] ~ dbin(%s[i,%d], %s[i,%d])", rv, k, pv, k, nv, k)))
    }

    # Logit link: arm 1 = baseline, subsequent arms += random effect
    out = paste0(out, L(sprintf("    logit(%s[i,1]) <- %s[i]", pv, muv)))

    if (na == 2) {
      out = paste0(out,
        L(sprintf("    logit(%s[i,2]) <- %s[i] + delta.2[i]", pv, muv)),
        L(sprintf("    delta.2[i] ~ dnorm(d[%s[i,2]] - d[%s[i,1]], prec.tau)", tv, tv))
      )
    } else {
      for (k in 2:na) {
        out = paste0(out, L(sprintf(
          "    logit(%s[i,%d]) <- %s[i] + %s[i,%d]", pv, k, muv, xiv, k - 1)))
      }
      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 (identical structure to continuous)
  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)
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)

# Frequentist comparison (log-OR via netmeta)
meta::pairwise(treat = treat, event = r, n = n,
               studlab = study, data = arm.stats, sm = "OR") %>%
  netmeta::netmetabin(event1, n1, event2, n2, treat1, treat2,
                      studlab, data = ., sm = "OR",
                      method = "LRP")


# 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 (d[t] = log-ORs vs reference) ---------------------------

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 log-ORs vs reference (d[t]) ---\n")
ps.d = posterior.summary(post.mat, d.params)
rownames(ps.d) = trt.labels
print(round(ps.d, 3))

# OR scale (exponentiate posterior draws)
cat("\n--- Pooled ORs vs reference (exp(d[t])) ---\n")
or.mat = exp(post.mat[, d.params, drop = FALSE])
colnames(or.mat) = d.params
ps.or  = posterior.summary(or.mat, d.params)
rownames(ps.or) = trt.labels
print(round(ps.or, 3))

cat("\n--- Between-study heterogeneity (tau, log-OR scale) ---\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. League table (log-OR scale; exponentiate cell strings for OR scale) ------

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))
  ) %>%
  # Cell shows log-OR [95% HDI]; swap exp() calls below for OR scale
  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)) %>%
  mutate(across(-t, ~ifelse(is.na(.), "\u2014", .))) %>%
  arrange(t) %>%
  rename(` ` = t)

cat("\n--- League table (log-OR [95% HDI], row vs column) ---\n")
print(league, row.names = FALSE)
