## ----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)
}

## ----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(is.na(x), "centralized",
    ifelse(x == 0, "$0$",
    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, "$0$",
    ifelse(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_rho_sweep <- function(tab)
    ktab(tab, row.names = FALSE,
                 col.names = c("$\\rho$", "Iterations", "Converged"),
                 caption = "Consensus-ADMM convergence on the surrogate cohort")
show_summary <- function(tab) {
    tab$sigma   <- tex_pow10(tab$sigma)
    tab$max_dev <- tex_sci(tab$max_dev)
    ktab(tab, digits = 6, row.names = FALSE, align = "lrrrrr",
                 col.names = c("$\\sigma$", "intercept", "age", "bmi", "sex",
                               "$\\max_j \\lvert \\hat z_j - \\hat\\beta_j^{\\text{centralized}} \\rvert$"),
                 caption = "DP-ADMM coefficients vs the centralized CVXR fit")
}
show_budget <- function(tab, T_iter) {
    tab$sigma <- tex_pow10(tab$sigma)
    for (v in c("rho_total", "epsilon_at_delta_1e_minus_5"))
        tab[[v]] <- tex_sci(tab[[v]])
    ktab(tab, row.names = FALSE, align = "lrr",
                 col.names = c("$\\sigma$",
                               sprintf("$\\rho_{\\text{total}} = %d\\,\\rho$", T_iter),
                               "$\\varepsilon$ at $\\delta = 10^{-5}$"),
                 caption = "zCDP composition; sensitivity $\\Delta = 1$, target $\\delta = 10^{-5}$")
}

## -----------------------------------------------------------------------------
suppressPackageStartupMessages({
    library(homomorpheR)
    library(CVXR)
    library(S7)
})

N   <- 3L
p   <- 4L
lam <- 1

## -----------------------------------------------------------------------------
build_local_problem <- function(X_i, y_i, rho_val) {
    x  <- Variable(p)
    zp <- Parameter(p)
    up <- Parameter(p)
    y_signs <- 2 * y_i - 1
    margins <- -y_signs * (X_i %*% x)
    local_loss <- sum(logistic(margins)) +
                  (lam / (2 * N)) * sum_squares(x)
    augmented  <- (rho_val / 2) * sum_squares(x - zp + up)
    prob <- Problem(Minimize(local_loss + augmented))
    value(zp) <- rep(0, p); value(up) <- rep(0, p)
    list(prob = prob, x = x, zp = zp, up = up)
}

## Inherits homomorpheR's abstract `Site` (which supplies `name` and the
## `state` environment), so it can take part in threshold key generation
## and keep its own share.
ConsensusSite <- new_class("ConsensusSite",
    parent     = homomorpheR::Site,
    properties = list(n = class_integer))

make_consensus_site <- function(name, X_i, y_i, rho_val) {
    st        <- new.env(parent = emptyenv())
    st$X      <- X_i
    st$y      <- y_i
    built     <- build_local_problem(X_i, y_i, rho_val)
    st$prob   <- built$prob
    st$x_var  <- built$x
    st$zp     <- built$zp
    st$up     <- built$up
    st$x_curr <- rep(0, ncol(X_i))
    st$u_curr <- rep(0, ncol(X_i))
    ConsensusSite(name = name, n = nrow(X_i), state = st)
}

local_update <- function(site, z_curr) {
    st <- site@state
    value(st$zp) <- z_curr
    value(st$up) <- st$u_curr
    suppressMessages(suppressWarnings(psolve(st$prob, solver = "CLARABEL")))
    if (!status(st$prob) %in% c("optimal", "optimal_inaccurate"))
        stop("Local CVXR solve at ", site@name, " did not reach optimal status.")
    st$x_curr <- as.numeric(value(st$x_var))
    invisible(st$x_curr)
}

## -----------------------------------------------------------------------------
set.seed(20260412)
n_per_site <- c(500L, 1000L, 1500L)
beta_true  <- c(intercept = -0.5, age = 0.4, bmi = -0.3, sex = 0.6)

make_site_data <- function(n) {
    X  <- cbind(1, rnorm(n), rnorm(n), rbinom(n, 1, 0.5))
    pr <- plogis(as.numeric(X %*% beta_true))
    y  <- as.integer(runif(n) < pr)
    list(X = X, y = y)
}
site_data <- lapply(n_per_site, make_site_data)

## ----eval=RECOMPUTE-----------------------------------------------------------
# tol      <- 1e-3
# max_iter <- 60L
# 
# ## Public design facts: three sites of these sizes, four
# ## covariates of these types. Nominal effect sizes, not the cohort's.
# beta_nominal   <- c(0, 0.5, 0.5, 0.5)
# surrogate_seed <- 20260413L
# 
# set.seed(surrogate_seed)
# surrogate_data <- lapply(n_per_site, function(n) {
#     X  <- cbind(1, rnorm(n), rnorm(n), rbinom(n, 1, 0.5))
#     pr <- plogis(as.numeric(X %*% beta_nominal))
#     list(X = X, y = as.integer(runif(n) < pr))
# })
# 
# sweep_one_rho <- function(cohort, rho_val) {
#     built <- lapply(cohort,
#                     function(s) build_local_problem(s$X, s$y, rho_val))
#     x_curr <- u_curr <- replicate(N, rep(0, p), simplify = FALSE)
#     z      <- rep(0, p)
#     k_conv <- NA_integer_
#     for (k in seq_len(max_iter)) {
#         for (i in seq_len(N)) {
#             value(built[[i]]$zp) <- z
#             value(built[[i]]$up) <- u_curr[[i]]
#             suppressMessages(suppressWarnings(
#                 psolve(built[[i]]$prob, solver = "CLARABEL")))
#             x_curr[[i]] <- as.numeric(value(built[[i]]$x))
#         }
#         z_prev <- z
#         z <- Reduce(`+`, Map(`+`, x_curr, u_curr)) / N
#         for (i in seq_len(N)) u_curr[[i]] <- u_curr[[i]] + x_curr[[i]] - z
#         pri <- sqrt(sum(vapply(seq_len(N),
#             function(i) sum((x_curr[[i]] - z)^2), 0)) / N)
#         dua <- rho_val * sqrt(sum((z - z_prev)^2))
#         if (pri < tol && dua < tol) { k_conv <- k; break }
#     }
#     data.frame(rho = rho_val,
#                iters = if (is.na(k_conv)) max_iter else k_conv,
#                converged = !is.na(k_conv))
# }
# 
# rho_grid  <- c(10, 20, 50, 100, 500)
# rho_sweep <- do.call(rbind,
#                      lapply(rho_grid,
#                             function(r) sweep_one_rho(surrogate_data, r)))
# show_rho_sweep(rho_sweep)
# 
# converged_rows <- rho_sweep[rho_sweep$converged, ]
# if (nrow(converged_rows) == 0L)
#     stop("No rho in the grid converged within max_iter on the surrogate.")
# 
# rho_chosen <- converged_rows$rho[which.min(converged_rows$iters)]
# T_fixed    <- converged_rows$iters[converged_rows$rho == rho_chosen]

## ----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))

## -----------------------------------------------------------------------------
## Site-side: the site draws its own noise, adds it, and encrypts with
## the public parameters it received at setup, all before anything
## leaves the site. The noiseless x_i + u_i never leaves.
site_contribution_dp <- function(site, sigma, Nv) {
    st <- site@state
    noised <- st$x_curr + st$u_curr + rnorm(p, mean = 0, sd = sigma * sqrt(Nv))
    encrypt(site, noised)
}

## Aggregator-side: sum the encrypted values, scale, threshold-decrypt. The
## 1/N scaling contracts the summed noise variance back to sigma^2.
encrypted_consensus_dp <- function(threshold_master, sites, sigma) {
    Nv  <- length(sites)
    cts <- lapply(sites, site_contribution_dp, sigma = sigma, Nv = Nv)
    ct_avg <- Reduce(`+`, cts) * (1 / Nv)
    decrypt(threshold_master, ct_avg, len = p)
}

## -----------------------------------------------------------------------------
run_dp_admm <- function(sigma, T_iter = T_fixed, seed = NULL) {
    if (!is.null(seed)) set.seed(seed)
    ## The sites exist first: the joint public key is built from them,
    ## each keeping the share it generates.
    sites <- list(
        make_consensus_site("Site 1", site_data[[1]]$X, site_data[[1]]$y, rho_chosen),
        make_consensus_site("Site 2", site_data[[2]]$X, site_data[[2]]$y, rho_chosen),
        make_consensus_site("Site 3", site_data[[3]]$X, site_data[[3]]$y, rho_chosen))
    master <- make_threshold_master("Aggregator",
                                    crypto_context = cc,
                                    sites          = sites)

    z_curr <- rep(0, p)
    z_hist <- matrix(NA_real_, nrow = T_iter, ncol = p,
                     dimnames = list(NULL, names(beta_true)))
    for (k in seq_len(T_iter)) {
        for (s in sites) local_update(s, z_curr)
        z_curr <- encrypted_consensus_dp(master, sites, sigma)
        for (s in sites) {
            s@state$u_curr <- s@state$u_curr + (s@state$x_curr - z_curr)
        }
        z_hist[k, ] <- z_curr
    }
    list(z = z_curr, z_hist = z_hist)
}

## -----------------------------------------------------------------------------
X_pooled  <- do.call(rbind, lapply(site_data, `[[`, "X"))
y_pooled  <- unlist(lapply(site_data, `[[`, "y"))
beta_var  <- Variable(p)
y_signs_p <- 2 * y_pooled - 1
margins_p <- -y_signs_p * (X_pooled %*% beta_var)
suppressMessages(suppressWarnings(
    psolve(Problem(Minimize(sum(logistic(margins_p)) +
                            (lam / 2) * sum_squares(beta_var))),
           solver = "CLARABEL")))
beta_central <- as.numeric(value(beta_var))
names(beta_central) <- names(beta_true)

## ----eval=RECOMPUTE-----------------------------------------------------------
# sigma_grid    <- c(0, 1e-4, 1e-3, 1e-2, 1e-1, 1)
# sweep_results <- vector("list", length(sigma_grid))
# for (j in seq_along(sigma_grid)) {
#     sweep_results[[j]] <- run_dp_admm(sigma = sigma_grid[j], seed = 100L + j)
# }
# names(sweep_results) <- sprintf("sigma=%.0e", sigma_grid)

## ----eval=RECOMPUTE-----------------------------------------------------------
# clean_dev <- max(abs(sweep_results[[1]]$z - beta_central))
# agree_tol <- 10 * tol
# if (clean_dev > agree_tol)
#     stop("DP-ADMM at sigma = 0 disagrees with the centralized fit.")

## ----eval=RECOMPUTE-----------------------------------------------------------
# summary_df <- do.call(rbind, lapply(seq_along(sigma_grid), function(j) {
#     z <- sweep_results[[j]]$z
#     data.frame(sigma     = sigma_grid[j],
#                intercept = z[1],
#                age       = z[2],
#                bmi       = z[3],
#                sex       = z[4],
#                max_dev   = max(abs(z - beta_central)))
# }))
# central_row <- data.frame(sigma = NA, intercept = beta_central[1],
#                          age = beta_central[2], bmi = beta_central[3],
#                          sex = beta_central[4], max_dev = 0)
# summary_table <- rbind(summary_df, central_row)
# rownames(summary_table) <- c(sprintf("sigma=%g", sigma_grid), "centralized")
# show_summary(summary_table)

## ----eval=RECOMPUTE, echo=FALSE-----------------------------------------------
# cvxr_admm_dp_results <- list(tol           = tol,
#                              rho_sweep     = rho_sweep,
#                              rho_chosen    = rho_chosen,
#                              T_fixed       = T_fixed,
#                              sigma_grid    = sigma_grid,
#                              clean_dev     = clean_dev,
#                              summary_table = summary_table)

## -----------------------------------------------------------------------------
zcdp_to_eps <- function(rho, delta = 1e-5) rho + 2 * sqrt(rho * log(1 / delta))

budget <- data.frame(sigma = sigma_grid[sigma_grid > 0])
budget$rho_total                  <- T_fixed * (1 / budget$sigma)^2 / 2
budget$epsilon_at_delta_1e_minus_5 <- zcdp_to_eps(budget$rho_total)
show_budget(budget, T_fixed)

