options(stringsAsFactors = FALSE)

input_file <- "HPV_OPMD_Meta_Analysis_Input.csv"
output_dir <- "analysis_results"

if (!dir.exists(output_dir)) {
  dir.create(output_dir, recursive = TRUE)
}

inv_logit <- function(x) {
  1 / (1 + exp(-x))
}

prepare_effects <- function(data, correction = TRUE) {
  events <- data$events
  non_events <- data$total - data$events

  if (correction && any(events == 0 | non_events == 0)) {
    events <- events + 0.5
    non_events <- non_events + 0.5
  }

  data$events_model <- events
  data$non_events_model <- non_events
  data$logit <- log(events / non_events)
  data$variance <- 1 / events + 1 / non_events
  data
}

fit_reml <- function(data, correction = TRUE) {
  d <- prepare_effects(data, correction)

  reml_objective <- function(tau2) {
    weights <- 1 / (d$variance + tau2)
    pooled <- sum(weights * d$logit) / sum(weights)
    0.5 * (
      sum(log(d$variance + tau2)) +
      log(sum(weights)) +
      sum(weights * (d$logit - pooled)^2)
    )
  }

  tau2 <- optimize(reml_objective, interval = c(0, 100), tol = 1e-12)$minimum
  weights <- 1 / (d$variance + tau2)
  pooled <- sum(weights * d$logit) / sum(weights)
  standard_error <- sqrt(1 / sum(weights))
  k <- nrow(d)

  ci_logit <- pooled + c(-1, 1) * qnorm(0.975) * standard_error
  pi_logit <- pooled + c(-1, 1) * qt(0.975, df = k - 1) *
    sqrt(tau2 + standard_error^2)

  fixed_weights <- 1 / d$variance
  fixed_pooled <- sum(fixed_weights * d$logit) / sum(fixed_weights)
  q_value <- sum(fixed_weights * (d$logit - fixed_pooled)^2)
  i2 <- max(0, (q_value - (k - 1)) / q_value) * 100

  q_hksj <- sum(weights * (d$logit - pooled)^2) / (k - 1)
  se_hksj <- sqrt(q_hksj / sum(weights))
  ci_hksj_logit <- pooled + c(-1, 1) * qt(0.975, df = k - 1) * se_hksj

  d$weight_percent <- 100 * weights / sum(weights)
  exact_ci <- t(vapply(
    seq_len(nrow(d)),
    function(i) binom.test(d$events[i], d$total[i])$conf.int,
    numeric(2)
  ))
  d$proportion <- d$events / d$total
  d$exact_ci_low <- exact_ci[, 1]
  d$exact_ci_high <- exact_ci[, 2]

  summary <- data.frame(
    k = k,
    events = sum(data$events),
    total = sum(data$total),
    pooled_prevalence = inv_logit(pooled),
    ci_low = inv_logit(ci_logit[1]),
    ci_high = inv_logit(ci_logit[2]),
    prediction_low = inv_logit(pi_logit[1]),
    prediction_high = inv_logit(pi_logit[2]),
    hksj_ci_low = inv_logit(ci_hksj_logit[1]),
    hksj_ci_high = inv_logit(ci_hksj_logit[2]),
    tau2 = tau2,
    q = q_value,
    q_df = k - 1,
    i2_percent = i2
  )

  list(studies = d, summary = summary)
}

data <- read.csv(input_file, check.names = FALSE)

primary <- fit_reml(data, correction = TRUE)
tissue <- fit_reml(subset(data, tissue_only == "Yes"), correction = TRUE)
without_zero <- fit_reml(subset(data, events > 0), correction = FALSE)
without_largest <- fit_reml(
  subset(data, !grepl("^Sundberg et al., 2021", study)),
  correction = TRUE
)

scenario_summary <- rbind(
  cbind(scenario = "Primary analysis", primary$summary),
  cbind(scenario = "Tissue/biopsy only", tissue$summary),
  cbind(scenario = "Exclusion of the zero-event study", without_zero$summary),
  cbind(scenario = "Exclusion of the largest study", without_largest$summary)
)

write.csv(
  primary$studies,
  file.path(output_dir, "study_level_results.csv"),
  row.names = FALSE
)
write.csv(
  scenario_summary,
  file.path(output_dir, "sensitivity_results.csv"),
  row.names = FALSE
)

png(
  file.path(output_dir, "forest_plot.png"),
  width = 1800,
  height = 1350,
  res = 180
)
par(mar = c(5, 12, 3, 2), family = "sans")
y <- rev(seq_len(nrow(primary$studies)))
plot(
  primary$studies$proportion,
  y,
  xlim = c(0, 1),
  ylim = c(0, nrow(primary$studies) + 2),
  pch = 15,
  xlab = "HPV prevalence",
  ylab = "",
  yaxt = "n",
  main = "Study-level HPV prevalence"
)
segments(
  primary$studies$exact_ci_low,
  y,
  primary$studies$exact_ci_high,
  y
)
axis(2, at = y, labels = primary$studies$study, las = 1, cex.axis = 0.75)
abline(v = primary$summary$pooled_prevalence, lty = 2)
points(primary$summary$pooled_prevalence, 0.7, pch = 18, cex = 1.4)
segments(
  primary$summary$ci_low,
  0.7,
  primary$summary$ci_high,
  0.7,
  lwd = 2
)
dev.off()

png(
  file.path(output_dir, "funnel_plot.png"),
  width = 1350,
  height = 1200,
  res = 180
)
par(mar = c(5, 5, 3, 2), family = "sans")
study_se <- sqrt(primary$studies$variance)
plot(
  primary$studies$logit,
  study_se,
  pch = 16,
  xlab = "Logit prevalence",
  ylab = "Standard error",
  main = "Funnel plot",
  ylim = rev(range(c(0, study_se)))
)
abline(v = qlogis(primary$summary$pooled_prevalence), lty = 2)
dev.off()

print(primary$summary, digits = 6)
print(scenario_summary, digits = 6)
