setwd("D:/charls/")
library(lme4)
library(lmerTest)
library(haven)
library(tidyverse)
library(lm.beta)
library(tableone)
library(charlsMAX)
library(httr)
base <- read.csv("base.csv")

#NDVI
model1_NDVI <- glm(
  blood ~ NDVI_Q, 
  family = binomial,
  data = base
)

model2_NDVI <- glm(
  blood ~ NDVI_Q + age + sex + marriage + house + activity, 
  family = binomial,
  data = base
)

model3_NDVI <- glm(
  blood ~ NDVI_Q + age + sex + marriage + house + activity + smoke + drink + health + sleep + depression + BMI, 
  family = binomial,
  data = base
)

#PM2.5
model1_PM2.5 <- glm(
  blood ~ PM2.5_Q, 
  family = binomial,
  data = base
)

model2_PM2.5 <- glm(
  blood ~ PM2.5_Q + age + sex + marriage + house + activity, 
  family = binomial,
  data = base
)

model3_PM2.5 <- glm(
  blood ~ PM2.5_Q + age + sex + marriage + house + activity + smoke + drink + health + sleep + depression + BMI, 
  family = binomial,
  data = base
)

#NO2
model1_NO2 <- glm(
  blood ~ NO2_Q, 
  family = binomial,
  data = base
)

model2_NO2 <- glm(
  blood ~ NO2_Q + age + sex + marriage + house + activity, 
  family = binomial,
  data = base
)

model3_NO2 <- glm(
  blood ~ NO2_Q + age + sex + marriage + house + activity + smoke + drink + health + sleep + depression + BMI, 
  family = binomial,
  data = base
)

#SO2
model1_SO2 <- glm(
  blood ~ SO2_Q, 
  family = binomial,
  data = base
)

model2_SO2 <- glm(
  blood ~ SO2_Q + age + sex + marriage + house + activity, 
  family = binomial,
  data = base
)

model3_SO2 <- glm(
  blood ~ SO2_Q + age + sex + marriage + house + activity + smoke + drink + health + sleep + depression + BMI, 
  family = binomial,
  data = base
)

extract_pollutant_results <- function(model, pollutant_var) {
  
  coef_summary <- summary(model)$coefficients
  coef_row <- coef_summary[pollutant_var, ]
  or_value <- exp(coef_row[1])
  ci_values <- exp(confint(model, parm = pollutant_var, level = 0.95))
  p_value <- coef_summary[pollutant_var, 4]
  if (p_value < 0.001) {
    p_formatted <- "<0.001"
  } else {
    p_formatted <- sprintf("%.3f", p_value)
  }
  results_df <- data.frame(
    Pollutant = gsub("_Q", "", pollutant_var),
    Model_Type = ifelse(length(model$coefficients) == 2, "Model 1 (Unadjusted)",
                        ifelse(length(model$coefficients) == 7, "Model 2 (Demographics adjusted)",
                               "Model 3 (Fully adjusted)")),
    OR = or_value,
    Lower_CI = ci_values[1],
    Upper_CI = ci_values[2],
    P_value = p_formatted,
    stringsAsFactors = FALSE
  )
  
  return(results_df)
}

all_pollutant_results <- list()

all_pollutant_results[[1]] <- extract_pollutant_results(model1_NDVI, "NDVI_Q")
all_pollutant_results[[2]] <- extract_pollutant_results(model2_NDVI, "NDVI_Q")
all_pollutant_results[[3]] <- extract_pollutant_results(model3_NDVI, "NDVI_Q")

all_pollutant_results[[4]] <- extract_pollutant_results(model1_PM2.5, "PM2.5_Q")
all_pollutant_results[[5]] <- extract_pollutant_results(model2_PM2.5, "PM2.5_Q")
all_pollutant_results[[6]] <- extract_pollutant_results(model3_PM2.5, "PM2.5_Q")

all_pollutant_results[[7]] <- extract_pollutant_results(model1_NO2, "NO2_Q")
all_pollutant_results[[8]] <- extract_pollutant_results(model2_NO2, "NO2_Q")
all_pollutant_results[[9]] <- extract_pollutant_results(model3_NO2, "NO2_Q")

all_pollutant_results[[10]] <- extract_pollutant_results(model1_SO2, "SO2_Q")
all_pollutant_results[[11]] <- extract_pollutant_results(model2_SO2, "SO2_Q")
all_pollutant_results[[12]] <- extract_pollutant_results(model3_SO2, "SO2_Q")

pollutant_results <- do.call(rbind, all_pollutant_results)

# 绘制森林图
library(tidyverse)
results <- tibble(
  Pollutant = rep(c("NDVI", "PM2.5","NO2","SO2"), each = 3),
  Model = rep(c("Model 1", "Model 2", "Model 3"), 4),
  OR = pollutant_results$OR,          
  Lower = pollutant_results$Lower_CI,       
  Upper = pollutant_results$Upper_CI,    
  P_value = pollutant_results$P_value,         
) %>%
  mutate(
    OR_CI = sprintf("%.3f (%.3f–%.3f)", OR, Lower, Upper),
    Model = factor(Model, levels = c("Model 3", "Model 2", "Model 1"))
  )

ggplot(results, aes(x = OR, y = Model, color = Model, shape = Model)) +
  geom_vline(xintercept = 1, linetype = "dashed", color = "gray50", linewidth = 0.5) +
  geom_point(position = position_dodge(width = 0.7), size = 3) +
  geom_errorbarh(
    aes(xmin = Lower, xmax = Upper),
    position = position_dodge(width = 0.7),
    height = 0.15,
    linewidth = 0.8
  ) +
  geom_text(aes(x = 1.7, label = OR_CI), hjust = 0, size = 3.5, color = "black") +
  geom_text(aes(x = 2.1, label = P_value), hjust = 0, size = 3.5, color = "black") +
  annotate(
    "text", 
    x = 1.77, 
    y = Inf, 
    label = "OR(95%CI)", 
    vjust = 0.95
  ) +
  annotate(
    "text", 
    x = 2.13, 
    y = Inf, 
    label = "P value", 
    vjust = 0.95
  ) +
  facet_grid(Pollutant ~ ., scales = "free_y", switch = "y") +
  scale_x_log10(
    breaks = c(0.7, 0.8, 0.9, 1.0, 1.2, 1.4, 1.6),
    expand = expansion(mult = c(0.05, 0.1))
  ) +
  scale_color_manual(
    values = c("#1f77b4", "#ff7f0e", "#2ca02c"),
    breaks = c("Model 1", "Model 2", "Model 3"), 
    labels = c(
      "Model 1: Unadjusted model",
      "Model 2: Model adjusted for demographic factors",
      "Model 3: Model further adjusted for health behaviors"
    )
  ) +
  scale_shape_manual(
    values = c(16, 17, 15),
    breaks = c("Model 1", "Model 2", "Model 3"),
    labels = c(
      "Model 1: Unadjusted model",
      "Model 2: Model adjusted for demographic factors",
      "Model 3: Model further adjusted for health behaviors"
    )
  ) +
  labs(
    x = "Odds Ratio (95% CI)", 
    y = "", 
    color = "",
    shape = "",
    caption = "IQR: NDVI = 1519.87, PM2.5 = 26.99 μg/m³, NO2 = 15.62 μg/m³, SO2 = 17.63 μg/m³"
  ) +
  theme_minimal(base_size = 12) +
  theme(
    plot.caption = element_text(
      hjust = 0,
      margin = margin(t = 5)
    ),
    strip.placement = "outside",
    strip.text.y.left = element_text(
      angle = 0,
      hjust = 0,
      face = "bold",
      size = 12
    ),
    panel.grid.major.y = element_blank(),
    panel.grid.minor = element_blank(),
    axis.line.x = element_line(color = "gray30"),
    axis.ticks.x = element_line(color = "gray30"),
    axis.text.y = element_text(face = "bold"),
    legend.text = element_text(size = 14, face = "plain"),
    legend.title = element_text(size = 14, face = "plain"),
    legend.key.width = unit(1, "cm"),
    legend.position = "top",
    panel.spacing = unit(1, "lines")
  )