# This script keeps the core analytical steps used in the manuscript:
# 1) main time-stratified case-crossover DLNM models for delta ea, delta AH, delta RH;
# 2) subgroup analyses for delta ea;
# 3) scenario-based future OR/AF summaries for delta ea;
# 4) sensitivity analyses for lag window, covariate adjustment, and linear pollutants.

required_packages <- c("dplyr", "survival", "dlnm", "splines", "lubridate")

require_analysis_packages <- function() {
  missing_packages <- required_packages[
    !vapply(required_packages, requireNamespace, logical(1), quietly = TRUE)
  ]
  if (length(missing_packages) > 0) {
    stop(
      "Missing required R packages: ", paste(missing_packages, collapse = ", "),
      ". Install them before running, e.g. install.packages(c(",
      paste(sprintf("'%s'", missing_packages), collapse = ", "), "))."
    )
  }
  suppressPackageStartupMessages({
    library(dplyr)
    library(survival)
    library(dlnm)
    library(splines)
    library(lubridate)
  })
}

config <- list(
  analysis_csv = "Clinical_exposure_rhinitis_30d_time_ea.csv",
  future_csv = "pop_weighted_eacn_daily_all_scenarios.csv",
  output_dir = "results_public",
  id_col = "patient_id",
  case_col = "case",
  date_col = "date",
  sex_col = "sex",
  age_col = "age",
  metrics = c(eacn = "eacn", AHCN = "AHCN", RHCN = "RHCN"),
  main_metric = "eacn",
  lag_max = 7,
  exposure_df = 3,
  lag_df = 3,
  covar_df = 3,
  percentiles = c(P5 = 0.05, P10 = 0.10, P90 = 0.90, P95 = 0.95),
  full_covars = c(
    "lag_0_t2m", "lag_0_RH", "lag_0_wind", "lag_0_sp",
    "lag_0_TCN", "lag_0_tp",
    "lag_0_PM2.5", "lag_0_O3", "lag_0_NO2", "lag_0_SO2"
  )
)

read_analysis_csv <- function(path) {
  if (!file.exists(path)) stop("Input file not found: ", path)
  utils::read.csv(path, check.names = FALSE, stringsAsFactors = FALSE)
}

write_result <- function(x, file) {
  dir.create(dirname(file), showWarnings = FALSE, recursive = TRUE)
  utils::write.csv(x, file, row.names = FALSE, fileEncoding = "UTF-8")
}

require_columns <- function(dat, cols) {
  missing_cols <- setdiff(cols, names(dat))
  if (length(missing_cols) > 0) {
    stop("Missing required columns: ", paste(missing_cols, collapse = ", "))
  }
}

lag_names <- function(metric, lag_max) {
  paste0("lag_", 0:lag_max, "_", metric)
}

to_numeric_columns <- function(dat, cols) {
  for (v in intersect(cols, names(dat))) dat[[v]] <- as.numeric(dat[[v]])
  dat
}

keep_informative_strata <- function(dat, id_col, case_col) {
  dat[[id_col]] <- as.factor(dat[[id_col]])
  ids_ok <- dat %>%
    group_by(.data[[id_col]]) %>%
    summarise(
      has_case = any(.data[[case_col]] == 1),
      has_ctrl = any(.data[[case_col]] == 0),
      .groups = "drop"
    ) %>%
    filter(has_case & has_ctrl) %>%
    pull(.data[[id_col]]) %>%
    unique()
  dat %>% filter(.data[[id_col]] %in% ids_ok)
}

covariate_terms <- function(covars, df = 3, linear_covars = character(0)) {
  if (length(covars) == 0) return(character(0))
  vapply(covars, function(v) {
    if (v %in% linear_covars) {
      paste0("`", v, "`")
    } else {
      sprintf("ns(`%s`, df=%d)", v, df)
    }
  }, character(1))
}

prepare_model_data <- function(dat, metric, lag_max, covars, cfg) {
  exposure_cols <- lag_names(metric, lag_max)
  needed <- unique(c(cfg$id_col, cfg$case_col, exposure_cols, covars, paste0("lag_0_", metric)))
  require_columns(dat, needed)

  dat <- to_numeric_columns(dat, c(exposure_cols, covars, paste0("lag_0_", metric)))
  dat <- keep_informative_strata(dat, cfg$id_col, cfg$case_col)

  exposure_matrix <- as.matrix(dat[, exposure_cols, drop = FALSE])
  covar_matrix <- if (length(covars) > 0) dat[, covars, drop = FALSE] else NULL
  rows_ok <- if (length(covars) > 0) {
    complete.cases(cbind(exposure_matrix, covar_matrix))
  } else {
    complete.cases(exposure_matrix)
  }

  model_cols <- unique(c(cfg$id_col, cfg$case_col, covars, paste0("lag_0_", metric)))
  list(
    data = dat[rows_ok, model_cols, drop = FALSE],
    exposure_matrix = exposure_matrix[rows_ok, , drop = FALSE]
  )
}

fit_dlnm_clogit <- function(dat, metric, lag_max, covars, cfg,
                            linear_covars = character(0)) {
  prepared <- prepare_model_data(dat, metric, lag_max, covars, cfg)
  model_dat <- prepared$data
  exposure_matrix <- prepared$exposure_matrix

  cb <- crossbasis(
    exposure_matrix,
    lag = c(0, lag_max),
    argvar = list(fun = "ns", df = cfg$exposure_df),
    arglag = list(fun = "ns", df = cfg$lag_df)
  )

  covar_part <- covariate_terms(covars, cfg$covar_df, linear_covars)
  rhs <- c("cb", covar_part, paste0("strata(`", cfg$id_col, "`)"))
  fml <- as.formula(paste0(cfg$case_col, " ~ ", paste(rhs, collapse = " + ")))
  fit <- clogit(fml, data = model_dat)

  list(
    fit = fit,
    crossbasis = cb,
    data = model_dat,
    exposure_matrix = exposure_matrix,
    metric = metric,
    lag_max = lag_max,
    covars = covars,
    aic = AIC(fit)
  )
}

get_cb_coef_vcov <- function(model, cb_prefix = "cb") {
  cf <- coef(model$fit)
  vc <- vcov(model$fit)
  cn <- names(cf)
  
  idx <- grepl(paste0("^", cb_prefix), cn)
  
  if (!any(idx)) {
    stop("Cannot find crossbasis coefficients in model.")
  }
  
  selected_coef <- cf[idx]
  selected_vcov <- vc[idx, idx, drop = FALSE]
  
  if (length(selected_coef) != ncol(model$crossbasis)) {
    stop(
      "Number of selected crossbasis coefficients does not match basis matrix. ",
      "Selected coefficients: ", length(selected_coef),
      "; basis columns: ", ncol(model$crossbasis)
    )
  }
  
  list(coef = selected_coef, vcov = selected_vcov)
}

predict_cumulative_curve <- function(model, center = 0, n_grid = 200) {
  rng <- range(model$exposure_matrix, na.rm = TRUE)
  if (!all(is.finite(rng)) || diff(rng) <= 0) {
    stop("Exposure has no finite variation for prediction.")
  }
  cv <- get_cb_coef_vcov(model)
  
  pred <- crosspred(
    model$crossbasis,
    coef = cv$coef,
    vcov = cv$vcov,
    cen = center,
    by = diff(rng) / n_grid,
    cumul = TRUE,
    model.link = "logit"
  )
  data.frame(
    metric = model$metric,
    lag_max = model$lag_max,
    exposure = pred$predvar,
    OR = pred$allRRfit,
    low = pred$allRRlow,
    high = pred$allRRhigh
  )
}

estimate_at_percentiles <- function(model, probs, cfg = config, center = 0) {
  lag0_col <- paste0("lag_0_", model$metric)
  case_exposure <- model$data[[lag0_col]][model$data[[cfg$case_col]] == 1]
  pct_values <- stats::quantile(case_exposure, probs = probs, na.rm = TRUE, names = TRUE)
  names(pct_values) <- names(probs)

  cv <- get_cb_coef_vcov(model)
  
  pred <- crosspred(
    model$crossbasis,
    coef = cv$coef,
    vcov = cv$vcov,
    cen = center,
    at = as.numeric(pct_values),
    cumul = TRUE,
    model.link = "logit"
  )

  data.frame(
    metric = model$metric,
    lag_max = model$lag_max,
    percentile = names(pct_values),
    threshold = as.numeric(pct_values),
    OR = pred$allRRfit,
    low = pred$allRRlow,
    high = pred$allRRhigh,
    AIC = model$aic
  )
}

run_main_models <- function(dat, cfg = config) {
  require_analysis_packages()
  out <- list()
  full_covars <- intersect(cfg$full_covars, names(dat))

  for (metric_name in names(cfg$metrics)) {
    metric <- unname(cfg$metrics[[metric_name]])
    model <- fit_dlnm_clogit(dat, metric, cfg$lag_max, full_covars, cfg)
    out[[metric]] <- list(
      model = model,
      percentile_OR = estimate_at_percentiles(model, cfg$percentiles, cfg),
      curve = predict_cumulative_curve(model)
    )
  }

  percentile_table <- bind_rows(lapply(out, `[[`, "percentile_OR"))
  curve_table <- bind_rows(lapply(out, `[[`, "curve"))
  write_result(percentile_table, file.path(cfg$output_dir, "main_percentile_OR.csv"))
  write_result(curve_table, file.path(cfg$output_dir, "main_cumulative_curves.csv"))
  out
}

add_subgroup_variables <- function(dat, cfg = config) {
  require_analysis_packages()
  ah_cols <- c("lag_0_AH", "lag_1_AH", "lag_2_AH")
  require_columns(dat, c(cfg$id_col, cfg$case_col, cfg$date_col, cfg$sex_col, cfg$age_col, ah_cols))

  dat <- to_numeric_columns(dat, c(cfg$age_col, ah_cols))
  dat[[cfg$date_col]] <- as.Date(dat[[cfg$date_col]])

  dat <- dat %>%
    mutate(bg_AH_case_day = rowMeans(across(all_of(ah_cols)), na.rm = FALSE))

  bg_by_id <- dat %>%
    filter(.data[[cfg$case_col]] == 1) %>%
    group_by(.data[[cfg$id_col]]) %>%
    summarise(bg_AH = mean(bg_AH_case_day, na.rm = TRUE), .groups = "drop") %>%
    filter(is.finite(bg_AH))

  median_bg <- median(bg_by_id$bg_AH, na.rm = TRUE)
  bg_by_id <- bg_by_id %>%
    mutate(background_AH = ifelse(bg_AH <= median_bg, "Dry", "Humid"))

  dat <- left_join(dat, bg_by_id[, c(cfg$id_col, "background_AH")], by = cfg$id_col)
  month_value <- as.integer(format(dat[[cfg$date_col]], "%m"))

  dat %>%
    mutate(
      sex_group = case_when(
        .data[[cfg$sex_col]] %in% c("Male", "male", "M", "1", "男") ~ "Male",
        .data[[cfg$sex_col]] %in% c("Female", "female", "F", "2", "女") ~ "Female",
        TRUE ~ as.character(.data[[cfg$sex_col]])
      ),
      age_group = ifelse(.data[[cfg$age_col]] < 6, "<6 years", "6-18 years"),
      season = ifelse(month_value %in% 5:10, "Warm season", "Cold season")
    )
}

run_subgroup_models <- function(dat, cfg = config) {
  dat <- add_subgroup_variables(dat, cfg)
  full_covars <- intersect(cfg$full_covars, names(dat))
  metric <- cfg$main_metric

  subgroup_defs <- list(
    Overall = rep(TRUE, nrow(dat)),
    Male = dat$sex_group == "Male",
    Female = dat$sex_group == "Female",
    Age_lt_6 = dat$age_group == "<6 years",
    Age_6_18 = dat$age_group == "6-18 years",
    Warm = dat$season == "Warm season",
    Cold = dat$season == "Cold season",
    Dry_background = dat$background_AH == "Dry",
    Humid_background = dat$background_AH == "Humid"
  )

  results <- list()
  for (nm in names(subgroup_defs)) {
    subdat <- dat[which(subgroup_defs[[nm]] %in% TRUE), , drop = FALSE]
    if (nrow(subdat) == 0) next
    fit <- try(fit_dlnm_clogit(subdat, metric, cfg$lag_max, full_covars, cfg), silent = TRUE)
    if (inherits(fit, "try-error")) next
    tab <- estimate_at_percentiles(fit, cfg$percentiles, cfg)
    tab$subgroup <- nm
    results[[nm]] <- tab
  }

  subgroup_table <- bind_rows(results) %>%
    select(subgroup, everything())
  write_result(subgroup_table, file.path(cfg$output_dir, "subgroup_delta_ea_percentile_OR.csv"))
  subgroup_table
}

run_lag_sensitivity <- function(dat, cfg = config, lag_values = c(3, 5, 7, 14)) {
  require_analysis_packages()
  full_covars <- intersect(cfg$full_covars, names(dat))
  results <- list()
  for (lag_max in lag_values) {
    exposure_cols <- lag_names(cfg$main_metric, lag_max)
    if (!all(exposure_cols %in% names(dat))) next
    fit <- fit_dlnm_clogit(dat, cfg$main_metric, lag_max, full_covars, cfg)
    results[[as.character(lag_max)]] <- estimate_at_percentiles(fit, cfg$percentiles, cfg)
  }
  out <- bind_rows(results)
  write_result(out, file.path(cfg$output_dir, "sensitivity_lag_windows_OR.csv"))
  out
}

run_covariate_sensitivity <- function(dat, cfg = config) {
  require_analysis_packages()
  schemes <- list(
    M0_strata_only = character(0),
    M1_T_RH_deltaT = c("lag_0_t2m", "lag_0_RH", "lag_0_TCN"),
    M2_meteorology = c("lag_0_t2m", "lag_0_RH", "lag_0_wind", "lag_0_sp", "lag_0_TCN", "lag_0_tp"),
    M3_full = cfg$full_covars
  )

  results <- list()
  for (scheme in names(schemes)) {
    covars <- intersect(schemes[[scheme]], names(dat))
    fit <- fit_dlnm_clogit(dat, cfg$main_metric, cfg$lag_max, covars, cfg)
    tab <- estimate_at_percentiles(fit, cfg$percentiles, cfg)
    tab$model <- scheme
    results[[scheme]] <- tab
  }

  out <- bind_rows(results) %>% select(model, everything())
  write_result(out, file.path(cfg$output_dir, "sensitivity_covariate_models_OR.csv"))
  out
}

run_linear_pollutant_sensitivity <- function(dat, cfg = config) {
  require_analysis_packages()
  full_covars <- intersect(cfg$full_covars, names(dat))
  pollutants <- intersect(c("lag_0_PM2.5", "lag_0_O3", "lag_0_NO2", "lag_0_SO2"), full_covars)

  fit_ns <- fit_dlnm_clogit(dat, cfg$main_metric, cfg$lag_max, full_covars, cfg)
  fit_linear <- fit_dlnm_clogit(
    dat, cfg$main_metric, cfg$lag_max, full_covars, cfg,
    linear_covars = pollutants
  )

  out <- bind_rows(
    estimate_at_percentiles(fit_ns, cfg$percentiles, cfg) %>% mutate(model = "pollutants_ns"),
    estimate_at_percentiles(fit_linear, cfg$percentiles, cfg) %>% mutate(model = "pollutants_linear")
  ) %>%
    select(model, everything())

  write_result(out, file.path(cfg$output_dir, "sensitivity_linear_pollutants_OR.csv"))
  out
}

run_future_projection <- function(main_model, future_csv, cfg = config,
                                  scenario_cols = c("eacn_ssp126", "eacn_ssp370", "eacn_ssp585"),
                                  region_pattern = NULL) {
  require_analysis_packages()
  future <- read_analysis_csv(future_csv)
  require_columns(future, c("date", scenario_cols))

  if (!is.null(region_pattern) && "region" %in% names(future)) {
    future <- future[grepl(region_pattern, future$region), , drop = FALSE]
  }

  future$date <- as.Date(parse_date_time(future$date, orders = c("Ymd", "Y-m-d", "Y/m/d")))
  future <- future %>% arrange(date)
  for (v in scenario_cols) future[[v]] <- as.numeric(future[[v]])

  curve <- predict_cumulative_curve(main_model)
  curve$logOR <- log(curve$OR)
  curve$logOR_se <- (log(curve$high) - log(curve$low)) / (2 * 1.96)
  f_OR <- approxfun(curve$exposure, curve$OR, rule = 2)
  f_logOR_se <- approxfun(curve$exposure, curve$logOR_se, rule = 2)

  daily <- future[, c("date", scenario_cols), drop = FALSE]
  for (v in scenario_cols) {
    suffix <- sub("^eacn_", "", v)
    OR <- f_OR(daily[[v]])
    se_logOR <- f_logOR_se(daily[[v]])
    daily[[paste0("OR_", suffix)]] <- OR
    daily[[paste0("ORse_", suffix)]] <- OR * se_logOR
    daily[[paste0("AF_", suffix)]] <- 1 - 1 / OR
    daily[[paste0("AFse_", suffix)]] <- (1 / OR) * se_logOR
  }

  thresholds <- stats::quantile(
    main_model$data[[paste0("lag_0_", cfg$main_metric)]][main_model$data[[cfg$case_col]] == 1],
    probs = c(P5 = 0.05, P10 = 0.10),
    na.rm = TRUE
  )

  annual <- daily %>%
    mutate(year = year(date)) %>%
    group_by(year) %>%
    summarise(across(starts_with("AF_"), ~ mean(.x, na.rm = TRUE), .names = "mean_{.col}"),
              across(all_of(scenario_cols), list(
                drying_P5_days = ~ sum(.x <= thresholds["P5"], na.rm = TRUE),
                drying_P10_days = ~ sum(.x <= thresholds["P10"], na.rm = TRUE)
              )),
              .groups = "drop")

  decadal <- daily %>%
    mutate(year = year(date), decade = paste0((year %/% 10) * 10, "s")) %>%
    group_by(decade) %>%
    summarise(across(starts_with("AF_"), ~ mean(.x, na.rm = TRUE), .names = "mean_{.col}"),
              across(all_of(scenario_cols), list(
                drying_P5_days = ~ mean(.x <= thresholds["P5"], na.rm = TRUE) * 365.25,
                drying_P10_days = ~ mean(.x <= thresholds["P10"], na.rm = TRUE) * 365.25
              )),
              .groups = "drop")

  write_result(daily, file.path(cfg$output_dir, "future_daily_OR_AF.csv"))
  write_result(annual, file.path(cfg$output_dir, "future_annual_AF_drying_days.csv"))
  write_result(decadal, file.path(cfg$output_dir, "future_decadal_AF_drying_days.csv"))
  list(daily = daily, annual = annual, decadal = decadal, thresholds = thresholds)
}

run_all <- function(cfg = config) {
  require_analysis_packages()
  dir.create(cfg$output_dir, showWarnings = FALSE, recursive = TRUE)
  dat <- read_analysis_csv(cfg$analysis_csv)

  main <- run_main_models(dat, cfg)
  subgroup <- run_subgroup_models(dat, cfg)
  lag_sens <- run_lag_sensitivity(dat, cfg)
  covar_sens <- run_covariate_sensitivity(dat, cfg)
  pollutant_sens <- run_linear_pollutant_sensitivity(dat, cfg)

  future <- NULL
  if (file.exists(cfg$future_csv)) {
    future <- run_future_projection(main[[cfg$main_metric]]$model, cfg$future_csv, cfg)
  }

  invisible(list(
    main = main,
    subgroup = subgroup,
    lag_sensitivity = lag_sens,
    covariate_sensitivity = covar_sens,
    pollutant_sensitivity = pollutant_sens,
    future = future
  ))
}
