# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #
#                                                                             #
#   03-meta-analysis-binary.R                                                 #
#   Bayesian AD meta-analysis – binary outcome (y dichotomised)               #
#   Stage 1: per-study log RR from baseline-adjusted Poisson GLM              #
#            (Zou 2004; logit fallback if non-convergent)                     #
#   Stage 2: random-effects pooling in JAGS (log-RR scale)                    #
#   Prior:   mu ~ N(0, 1e-6);  tau ~ half-normal(0, 0.5)                      #
#                                                                             #
# # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # # #


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

library(readxl)
library(dplyr)
library(purrr)
library(runjags)
library(coda)
library(HDInterval)
library(ggplot2)
library(stringr)


# 1. Determine global median for dichotomisation -----------------------------

path   = "_data/data.xlsx"
sheets = readxl::excel_sheets(path)

global.median.y = purrr::map_dbl(sheets, function(sh) {
  readxl::read_excel(path, sheet = sh) %>% pull(y) %>% median(na.rm = TRUE)
}) %>% median()


# 2. Stage 1: per-study log RR from baseline-adjusted Poisson GLM -----------

es.list = purrr::map_dfr(sheets, function(sh) {

  dat = readxl::read_excel(path, sheet = sh) %>%
    dplyr::mutate(
      treat.bin = as.numeric(as.numeric(factor(treat))>1),
      y.bin     = as.integer(y > global.median.y))
  cc = dat[complete.cases(dat[, c("y.bin", "treat.bin", "baseline")]), ]
  fit = tryCatch(
    glm(y.bin ~ treat.bin + baseline, family = poisson(link = "log"), data = cc),
    error   = function(e) NULL,
    warning = function(w) {
      suppressWarnings(glm(y.bin ~ treat.bin + baseline,
                           family = poisson(link = "log"), data = cc))
    }
  ); fallback = FALSE
  if (is.null(fit) || !fit$converged) {
    fallback = TRUE
    fit = glm(y.bin ~ treat.bin + baseline,
              family = binomial(link = "logit"), data = cc) }
  log.rr    = coef(fit)[["treat.bin"]]
  var.logrr = vcov(fit)["treat.bin", "treat.bin"]
  data.frame(study = sh, n = nrow(cc), log.rr = log.rr, var.logrr = var.logrr,
             se.logrr = sqrt(var.logrr),
             metric   = ifelse(fallback, "log-OR (fallback)", "log-RR"))
})

es.list
if (any(es.list$metric != "log-RR"))
  message("Note: some studies used log-OR fallback — see 'metric' column.")


# 3. Stage 2: Bayesian random-effects meta-analysis in JAGS -----------------

# Model:
#   log.rr[j]  ~ N(theta[j],  var.logrr[j])  — known sampling variance
#   theta[j]   ~ N(mu, tau^2)                 — random effects
#   mu         ~ N(0, 1e-6)                   — flat prior on pooled log-RR
#   tau        ~ half-normal(0, 0.5)           — weakly informative heterogeneity prior

# The posterior of exp(mu) gives the pooled Risk Ratio.

k = nrow(es.list)
data.jags = list(
  k     = k,
  es    = es.list$log.rr,
  prec  = 1 / es.list$var.logrr   # known per-study precisions
)

M = "
model {

  for (j in 1:k) {
    es[j]    ~ dnorm(theta[j], prec[j])       # likelihood: observed log-RR
    theta[j] ~ dnorm(mu, inv.tau.sq)           # random effect for study j
  }

  mu         ~ dnorm(0, 1.0E-6)               # flat prior on pooled log-RR
  tau        ~ dnorm(0, pow(0.5, -2)) T(0,)   # half-normal(0, 0.5) on heterogeneity
  inv.tau.sq <- pow(tau, -2)

  # Predictive distribution for a new (unobserved) study
  theta.new ~ dnorm(mu, inv.tau.sq)
  rr.new <- exp(theta.new)

  # Back-transform for convenience
  rr <- exp(mu)
}

#monitor# mu, rr, tau, rr.new
"

fit.jags = run.jags(
  model    = M,
  data     = data.jags,
  n.chains = 4,
  burnin   = 5000,
  sample   = 80000,
  thin     = 2,
  summarise = FALSE
)
summary(fit.jags)

# Convergence check
mcmc.mat = coda::as.mcmc(fit.jags)

cat("\n── Gelman-Rubin R-hat ─────────────────────────────────────────────────\n")
print(coda::gelman.diag(fit.jags, multivariate = FALSE))
cat("\n── Effective Sample Size ──────────────────────────────────────────────\n")
print(round(coda::effectiveSize(fit.jags)))

