library(pROC)
library(caret)
library(data.table)

# Load data and PRS
data_all <- fread("data_all_prepared.txt", data.table = FALSE)
prs_info <- fread("ukb676305_PRS_new.txt", data.table = FALSE)

# Merge PRS
data_with_prs <- merge(data_all, prs_info, by = "eid")
data_prs <- data_with_prs[!is.na(data_with_prs$PRS_Standard), ]

# Split cases and controls
cases <- data_prs[data_prs$SLE == "1", ]
controls <- data_prs[data_prs$SLE == "0", ]
case_num <- nrow(cases)

result_prs <- matrix(nrow = 1000, ncol = 16)
colnames(result_prs) <- c("Coefficient", "Pvalue", "OR", "OR_lower", "OR_upper",
                          "AUC", "AUC_lower", "AUC_upper", "Accuracy", "Accuracy_lower",
                          "Accuracy_upper", "Balanced_accuracy", "Precision", 
                          "Sensitivity", "Specificity", "F1")

set.seed(666)
seeds <- sample(100:20000, 1000)

for (i in 1:1000) {
  set.seed(seeds[i])
  control_indices <- sample(1:nrow(controls), case_num)
  controls_sub <- controls[control_indices, ]
  data_merged <- rbind(cases, controls_sub)
  
  indices <- sample(nrow(data_merged))
  train_idx <- indices[1:round(0.7 * nrow(data_merged))]
  test_idx <- indices[(round(0.7 * nrow(data_merged)) + 1):nrow(data_merged)]
  train_data <- data_merged[train_idx, ]
  test_data <- data_merged[test_idx, ]
  
  # Logistic regression with PRS_Standard + age + sex
  fit <- glm(SLE ~ PRS_Standard + age + sex, data = train_data, family = binomial())
  coef_sum <- summary(fit)$coefficients
  result_prs[i, 1] <- coef_sum["PRS_Standard", "Estimate"]
  result_prs[i, 2] <- coef_sum["PRS_Standard", "Pr(>|z|)"]
  
  OR <- exp(coef(fit))
  CI <- exp(confint(fit))
  result_prs[i, 3] <- OR["PRS_Standard"]
  result_prs[i, 4] <- CI["PRS_Standard", 1]
  result_prs[i, 5] <- CI["PRS_Standard", 2]
  
  prob <- predict(fit, newdata = test_data, type = "response")
  pred_class <- ifelse(prob > 0.5, "1", "0")
  
  cm <- confusionMatrix(as.factor(pred_class), test_data$SLE, positive = "1")
  result_prs[i, 9] <- cm$overall["Accuracy"]
  result_prs[i, 10] <- cm$overall["AccuracyLower"]
  result_prs[i, 11] <- cm$overall["AccuracyUpper"]
  result_prs[i, 12] <- cm$byClass["Balanced Accuracy"]
  result_prs[i, 13] <- cm$byClass["Precision"]
  result_prs[i, 14] <- cm$byClass["Recall"]
  result_prs[i, 15] <- cm$byClass["Specificity"]
  result_prs[i, 16] <- cm$byClass["F1"]
  
  roc_obj <- roc(test_data$SLE, prob, levels = c("0", "1"), direction = "<", quiet = TRUE)
  result_prs[i, 6] <- as.numeric(auc(roc_obj))
  ci_auc <- ci.auc(roc_obj, conf.level = 0.95)
  result_prs[i, 7] <- ci_auc[1]
  result_prs[i, 8] <- ci_auc[3]
  
  if (i %% 100 == 0) cat("Completed PRS iteration", i, "\n")
}

result_prs <- as.data.frame(result_prs)
fwrite(result_prs, "Logistic_regression_performance_PRS.txt", row.names = TRUE, col.names = TRUE, sep = "\t")