## ----ktab, echo=FALSE---------------------------------------------------------
## Tables through kableExtra. Math written as $...$ in headers, cells
## and captions becomes \( ... \) and `code` becomes <code>, because
## pandoc does not process math inside a raw HTML table.
ktab <- function(x, ..., col.names = names(x), caption = NULL) {
    tex <- function(s)
        gsub("`([^`]*)`", "<code>\\1</code>",
             gsub("\\$([^$]+)\\$", "\\\\(\\1\\\\)", s))
    chr <- vapply(x, is.character, logical(1))
    x[chr] <- lapply(x[chr], tex)
    tab <- knitr::kable(x, format = "html", escape = FALSE,
                        col.names = tex(col.names),
                        caption = if (!is.null(caption)) tex(caption), ...)
    kableExtra::kable_styling(tab, bootstrap_options = c("striped", "condensed"),
                              full_width = TRUE)
}

## -----------------------------------------------------------------------------
suppressPackageStartupMessages(library(survival))
library(homomorpheR)
library(stats4)
data(DLBCL)

cox_data <- split(
  DLBCL[, c("time", "status", "GCB_sig", "LN_sig",
            "Prolif_sig", "BMP6", "MHC2_sig", "Subgroup")],
  DLBCL$Subgroup)

agg_model <- coxph(Surv(time, status) ~ GCB_sig + LN_sig +
                       Prolif_sig + BMP6 + MHC2_sig +
                       strata(Subgroup),
                   data = DLBCL)
agg_coef <- coef(agg_model)

cph_control <- replace(coxph.control(), "iter.max", 0)

local_cox_nll <- function(data, beta) {
    fit <- tryCatch(
        coxph(Surv(time, status) ~ GCB_sig + LN_sig + Prolif_sig +
                  BMP6 + MHC2_sig,
              data    = data,
              init    = beta,
              control = cph_control),
        error = function(e) NULL)
    if (is.null(fit)) NA_real_ else -fit$loglik[1]
}

## ----eval=RECOMPUTE-----------------------------------------------------------
# cc <- openfhe.R::fhe_context("CKKS",
#                            multiplicative_depth = 1L,
#                            scaling_mod_size     = 59L,
#                            first_mod_size       = 60L,
#                            batch_size           = 8L,
#                            features             = c(openfhe.R::Feature$MULTIPARTY))
# 
# n_sites <- length(cox_data)
# 
# build_dp_workers <- function(sigma) {
#     lapply(names(cox_data), function(nm) {
#         make_worker(
#             nm,
#             data     = cox_data[[nm]],
#             contribution_fn = function(data, beta) {
#                 nll <- local_cox_nll(data, beta)
#                 if (is.na(nll)) return(NA_real_)
#                 nll + rnorm(1L, mean = 0, sd = sigma / sqrt(n_sites))
#             })
#     })
# }
# 
# fit_at_sigma <- function(sigma, method = "BFGS", seed = 1L) {
#     set.seed(seed)   # stabilize the DP-noise draws across runs
#     workers <- build_dp_workers(sigma)
#     master  <- make_threshold_master("Aggregator",
#                                      crypto_context = cc,
#                                      sites          = workers)
#     ## Every call is one decrypted release, including the calls optim()
#     ## makes for finite-difference gradients and mle() for the Hessian.
#     n_queries <- 0L
#     dp_nLL <- function(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig) {
#         n_queries <<- n_queries + 1L
#         master_aggregate(master, c(GCB_sig, LN_sig, Prolif_sig, BMP6, MHC2_sig))
#     }
#     fit <- stats4::mle(dp_nLL,
#                        start   = list(GCB_sig = 0, LN_sig = 0, Prolif_sig = 0,
#                                       BMP6    = 0, MHC2_sig = 0),
#                        method  = method,
#                        control = list(reltol = 1e-7))
#     list(fit = fit, n_queries = n_queries)
# }

## ----echo=FALSE---------------------------------------------------------------
## Table display in mathematical notation. The stored tables keep
## their plain column names; only what is printed changes.
tex_pow10 <- function(x)
    ifelse(x == 1, "$1$", sprintf("$10^{%d}$", as.integer(round(log10(x)))))
tex_sci <- function(x, digits = 4) {
    e <- floor(log10(abs(x)))
    ifelse(x == 0 | (e >= -2 & e < 4),
           sprintf("$%s$", trimws(formatC(x, format = "fg", digits = digits))),
           sprintf("$%s \\times 10^{%d}$",
                   trimws(formatC(x / 10^e, format = "fg", digits = digits)),
                   as.integer(e)))
}
show_clean <- function(tab) {
    tab$abs_diff <- tex_sci(tab$abs_diff)
    ktab(tab, digits = 9, col.names = clean_cols, caption = clean_caption,
                 align = "lrrr")
}
clean_cols  <- c("Coefficient", "$\\hat\\beta$, `coxph()`",
                 "$\\hat\\beta$, protocol at $\\sigma = 0$",
                 "$\\lvert \\text{difference} \\rvert$")
clean_caption <- "Threshold-DP protocol at $\\sigma = 0$ vs cleartext `coxph()`"
sweep_cols  <- c("$\\sigma$", "Queries","GCB_sig", "LN_sig", "Prolif_sig",
                 "BMP6", "MHC2_sig",
                 "$\\max_j \\lvert \\hat\\beta_j - \\hat\\beta_j^{\\text{coxph}} \\rvert$")
budget_cols <- c("Optimizer", "$\\sigma$", "Queries $k$", "$\\rho$ per query",
                 "$\\rho_{\\text{total}} = k\\rho$", "$\\varepsilon$ at $\\delta = 10^{-5}$")
budget_caption <- "zCDP composition; sensitivity $\\Delta = 1$, target $\\delta = 10^{-5}$"
show_sweep <- function(tab, optimizer) {
    tab$sigma <- tex_pow10(tab$sigma)
    ktab(tab, digits = 6, row.names = FALSE, col.names = sweep_cols,
                 caption = sprintf("%s over the threshold-DP nLL at %d values of $\\sigma$",
                                   optimizer, nrow(tab)))
}
show_budget <- function(tab) {
    tab$sigma <- tex_pow10(tab$sigma)
    for (v in c("rho_per_query", "rho_total", "epsilon_at_delta_1e_minus_5"))
        tab[[v]] <- tex_sci(tab[[v]])
    ktab(tab, row.names = FALSE, col.names = budget_cols,
                 caption = budget_caption, align = "lrrrrr")
}

## ----eval=RECOMPUTE-----------------------------------------------------------
# fit_clean <- fit_at_sigma(0)$fit
# clean_check <- data.frame(
#     coefficient = names(agg_coef),
#     cleartext   = unname(agg_coef),
#     protocol    = unname(coef(fit_clean)[names(agg_coef)]),
#     abs_diff    = abs(unname(coef(fit_clean)[names(agg_coef)] - agg_coef))
# )
# show_clean(clean_check)

## ----eval=RECOMPUTE-----------------------------------------------------------
# sigma_grid <- c(1e-5, 1e-4, 1e-3, 1e-2, 1e-1, 1)
# 
# sweep_table <- function(fits) data.frame(
#     sigma        = sigma_grid,
#     n_queries    = sapply(fits, `[[`, "n_queries"),
#     GCB_sig      = sapply(fits, function(f) coef(f$fit)[["GCB_sig"]]),
#     LN_sig       = sapply(fits, function(f) coef(f$fit)[["LN_sig"]]),
#     Prolif_sig   = sapply(fits, function(f) coef(f$fit)[["Prolif_sig"]]),
#     BMP6         = sapply(fits, function(f) coef(f$fit)[["BMP6"]]),
#     MHC2_sig     = sapply(fits, function(f) coef(f$fit)[["MHC2_sig"]]),
#     max_abs_diff = sapply(fits, function(f)
#         max(abs(coef(f$fit) - agg_coef[names(coef(f$fit))]))))
# 
# fits_bfgs  <- lapply(sigma_grid, fit_at_sigma, method = "BFGS")
# bfgs_table <- sweep_table(fits_bfgs)
# show_sweep(bfgs_table, "BFGS")

## ----eval=RECOMPUTE-----------------------------------------------------------
# fits_nm  <- lapply(sigma_grid, fit_at_sigma, method = "Nelder-Mead")
# nm_table <- sweep_table(fits_nm)
# show_sweep(nm_table, "Nelder–Mead")

## ----eval=RECOMPUTE-----------------------------------------------------------
# zcdp_to_eps <- function(rho, delta = 1e-5) rho + 2 * sqrt(rho * log(1 / delta))
# 
# budget_rows <- function(optimizer, tab) data.frame(
#     optimizer = optimizer,
#     sigma     = tab$sigma[1:3],
#     n_queries = tab$n_queries[1:3])
# budget <- rbind(budget_rows("BFGS", bfgs_table),
#                 budget_rows("Nelder–Mead", nm_table))
# budget$rho_per_query              <- (1 / budget$sigma)^2 / 2
# budget$rho_total                  <- budget$n_queries * budget$rho_per_query
# budget$epsilon_at_delta_1e_minus_5 <- zcdp_to_eps(budget$rho_total)
# show_budget(budget)

## ----eval=RECOMPUTE, echo=FALSE-----------------------------------------------
# cox_threshold_dp_results <- list(clean_check = clean_check,
#                                  bfgs_table  = bfgs_table,
#                                  nm_table    = nm_table,
#                                  budget      = budget)

