# Fits an association's model on a set of cycles and reports the exposure's coefficient.

FAMILIES <- list(
  logistic = list(weighted = quote(quasibinomial()), unweighted = quote(binomial()), ratio = TRUE),
  linear = list(weighted = quote(gaussian()), unweighted = quote(gaussian()), ratio = FALSE),
  poisson = list(weighted = quote(quasipoisson()), unweighted = quote(poisson()), ratio = TRUE),
  log_binomial = list(weighted = quote(quasibinomial(link = "log")), unweighted = quote(binomial(link = "log")), ratio = TRUE)
)

# Every cycle's analysis frame, pooled, with the analysis weight w. An association builds each
# cycle's raw variables (build), may fix constants from its paper's own cycles (constants: such as
# quartile cutpoints or a standard deviation, used unchanged in 2021-2023), and then derives the
# model's variables (derive). In 2021-2023 the weight follows NCHS's directions for the cycle's
# subsamples (design.R) unless a$cycle_weights is FALSE, which keeps the paper's own weight (the
# sensitivity analysis); whether the supplement questionnaire is in the analysis is read from the
# components its build used.
pooled_data <- function(a, cycles) {
  frames <- lapply(cycles, function(cycle) {
    .read_log$components <- character()
    d <- a$build(cycle)
    if (!"SEQN" %in% names(d) || anyDuplicated(d$SEQN)) stop(a$id, ": build must give one row per SEQN")
    specific <- cycle == REPLICATION_CYCLE && !isFALSE(a$cycle_weights)
    blood_file <- if (specific) a$blood_file else NULL
    supplements <- specific && any(SUPPLEMENT_COMPONENTS %in% .read_log$components)
    w <- weights_for(a$weight, cycle, blood_file = blood_file, file = a$weight_file, supplements = supplements)
    d <- merge(d, w, by = "SEQN", all.x = TRUE)
    d$cycle <- cycle
    d$weight_name <- weight_variable(a$weight, cycle, blood = !is.null(blood_file), supplements = supplements)
    d
  })
  keep <- Reduce(intersect, lapply(frames, names))
  data <- do.call(rbind, lapply(frames, function(d) d[, keep]))
  data$w <- pooled_weight(data, a$weight, cycles, rule = if (is.null(a$weight_rule)) "nchs" else a$weight_rule)
  data
}

model_data <- function(a, cycles, constants = NULL) {
  data <- pooled_data(a, cycles)
  if (is.function(a$derive)) data <- a$derive(data, constants)
  data
}

# An association has two versions. "paper" is its paper's model and population, fitted on the
# paper's own cycles: the coding check. "harmonized" leaves out the covariates and exclusion steps
# that 2021-2023 can't build (a$formula_harmonized, the column in_population_harmonized), and is
# fitted on both the paper's cycles and 2021-2023, so what leaving them out changes is measured.
version_formula <- function(a, version) {
  if (version == "harmonized" && !is.null(a$formula_harmonized)) a$formula_harmonized else a$formula
}

version_population <- function(data, version) {
  if (version == "harmonized" && "in_population_harmonized" %in% names(data)) data$in_population_harmonized else data$in_population
}

# The coefficient of a$term, with its standard error, design degrees of freedom, p-value, and
# 95% confidence interval (t on the design's degrees of freedom).
fit_association <- function(a, cycles, constants = NULL, version = "paper") {
  data <- model_data(a, cycles, constants)
  formula <- version_formula(a, version)
  # Papers fit every model on one analytic sample: those with complete data for the full model.
  # A variant with fewer covariates (a$sample = "own" aside) keeps that sample.
  full <- if (!is.null(a$sample_formula)) a$sample_formula else formula
  variables <- union(all.vars(formula), if (identical(a$sample, "own")) character() else all.vars(full))
  missing <- setdiff(variables, names(data))
  if (length(missing) > 0) stop(a$id, ": no column ", paste(missing, collapse = ", "))
  family <- FAMILIES[[a$family]]
  weighted <- if (is.null(a$weighted)) TRUE else a$weighted
  # A weighted analysis covers those with a positive weight; an unweighted one doesn't need one.
  has_weight <- if (weighted) !is.na(data$w) & data$w > 0 else TRUE
  data$in_sample <- version_population(data, version) %in% TRUE & stats::complete.cases(data[, variables]) & has_weight
  # A covariate with a single value in the analytic sample can't be estimated (a cycle may lack a
  # category the paper's cycles had), so it is left out and reported, as the plan sets. The
  # exposure's own term is never left out.
  exposure <- function(label) startsWith(a$term, label)
  dropped <- Filter(function(label) {
    x <- data[[label]]
    !exposure(label) && !is.null(x) && length(unique(x[data$in_sample])) < 2
  }, attr(stats::terms(formula), "term.labels"))
  if (length(dropped) > 0) formula <- stats::update(formula, stats::as.formula(paste(". ~ .", paste("-", dropped, collapse = " "))))
  if (weighted) {
    design <- subset(survey_design(data), in_sample)
    fit <- survey::svyglm(formula, design = design, family = eval(family$weighted))
  } else {
    fit <- stats::glm(formula, data = data[data$in_sample, ], family = eval(family$unweighted))
  }
  coefficients <- summary(fit)$coefficients
  if (!a$term %in% rownames(coefficients)) stop(a$id, ": the model has no term ", a$term)
  b <- coefficients[a$term, 1]
  se <- coefficients[a$term, 2]
  # Design degrees of freedom: PSUs minus strata, as NCHS's guidelines use for t tests, not reduced
  # by the number of coefficients (survey's df.residual), which a model with many covariates can
  # drive to zero.
  df <- if (weighted) survey::degf(design) else Inf
  t <- b / se
  p <- 2 * stats::pt(-abs(t), df)
  half <- stats::qt(0.975, df) * se
  sample <- data[data$in_sample, ]
  y <- sample[[all.vars(formula)[1]]]
  list(
    b = b, se = se, df = df, p = p, ci_b = c(b - half, b + half),
    estimate = if (family$ratio) exp(b) else b,
    ci = if (family$ratio) exp(c(b - half, b + half)) else c(b - half, b + half),
    n = nrow(sample),
    events = if (a$family %in% c("logistic", "log_binomial")) sum(y == 1) else NA,
    weighted = weighted,
    weights = if (weighted) sort(unique(sample$weight_name)) else NULL,
    dropped = dropped,
    psu = if (weighted) nrow(unique(sample[, c("SDMVSTRA", "SDMVPSU")])) else NA,
    strata = if (weighted) length(unique(sample$SDMVSTRA)) else NA
  )
}
