#!/usr/bin/env Rscript
suppressPackageStartupMessages({
  library(RUVSeq)
  library(edgeR)
  library(ggplot2)
  library(ggrepel)
})

# ============================================================
# 0. CONFIG
# ============================================================
outdir  <- "/BRC/yan/heart/analysis/sn_ATAC_analysis/"
infile  <- paste0(outdir, "count_matrix_snPeak.txt")
k_range <- 1:6

cpm_cutoff <- 1
fdr_cutoff <- 0.001
fc_cutoff  <- log2(1.5)        # |log2FC| > log2(1.5)

# ============================================================
# 1. LOAD DATA
# ============================================================
cat("[1] Loading count matrix...\n")
counts           <- read.table(infile, header = TRUE, sep = "\t")
rownames(counts) <- paste0(counts$chr, ":", counts$start, "-", counts$end)
counts           <- counts[, -(1:3)]
cat("    Peaks:", nrow(counts), " Samples:", ncol(counts), "\n")

# ============================================================
# 2. METADATA
# ============================================================
cat("[2] Building metadata...\n")
meta <- data.frame(
  sample = colnames(counts),
  sex    = ifelse(grepl("female", colnames(counts)), "female", "male"),
  row.names = colnames(counts)
)
x      <- factor(meta$sex, levels = c("male", "female"))
colors <- ifelse(meta$sex == "female", "#E64B35", "#4DBBD5")
cat("    female:", sum(meta$sex == "female"),
    " male:",   sum(meta$sex == "male"), "\n")

# ============================================================
# 3. FILTER — CPM > 1 in >= min(group size) samples
# ============================================================
cat("[3] Filtering low-count peaks (CPM >", cpm_cutoff, ")...\n")
dge     <- DGEList(counts = as.matrix(counts), group = x)
cpm_mat <- cpm(counts)
keep    <- rowSums(cpm_mat > cpm_cutoff) >= min(table(x))
dge     <- dge[keep, , keep.lib.sizes = FALSE]
cpm_mat <- cpm_mat[keep, ]
cat("    Peaks after filtering:", nrow(dge), "\n")

# ============================================================
# 4. PRE-BATCH NORMALIZE + PCA
# ============================================================
cat("[4] Pre-batch normalization and PCA...\n")
dge <- calcNormFactors(dge)   # TMM normalization

log_pre <- log1p(t(cpm(dge)))
pca_pre <- prcomp(log_pre, scale. = FALSE)
pct_pre <- round(summary(pca_pre)$importance[2, 1:2] * 100, 1)
df_pre  <- data.frame(
  PC1    = pca_pre$x[, 1],
  PC2    = pca_pre$x[, 2],
  sex    = meta[rownames(pca_pre$x), "sex"],
  sample = rownames(pca_pre$x)
)
p_pre <- ggplot(df_pre, aes(PC1, PC2, color = sex, label = sample)) +
  geom_point(size = 3) +
  geom_text_repel(size = 2.5, max.overlaps = 20) +
  scale_color_manual(values = c(female = "#E64B35", male = "#4DBBD5")) +
  labs(title = "Heart snATAC PCA (pre-batch)",
       x = paste0("PC1 (", pct_pre[1], "%)"),
       y = paste0("PC2 (", pct_pre[2], "%)")) +
  theme_bw() +
  theme(plot.title = element_text(hjust = 0.5, face = "bold"))
ggsave(paste0(outdir, "heart_PCA_pre.pdf"), p_pre, width = 8, height = 6)

# ============================================================
# 5. RUVr RESIDUALS
# ============================================================
cat("[5] Computing deviance residuals for RUVr...\n")
set    <- newSeqExpressionSet(as.matrix(dge$counts),
           phenoData = data.frame(x, row.names = colnames(dge)))
design <- model.matrix(~ x, data = pData(set))

y_tmp  <- DGEList(counts = counts(set), group = x)
y_tmp  <- calcNormFactors(y_tmp, method = "upperquartile")
y_tmp  <- estimateGLMCommonDisp(y_tmp, design)
y_tmp  <- estimateGLMTagwiseDisp(y_tmp, design)
fit    <- glmFit(y_tmp, design)
res    <- residuals(fit, type = "deviance")
seqUQ  <- betweenLaneNormalization(set, which = "upper")

# ============================================================
# 6. RUVr k=1~6: diagnostics + DAR per k
# ============================================================
cat("[6] Running RUVr k=1~6 with full DAR each...\n")

results_summary <- data.frame()
dar_lists       <- list(all  = vector("list", length(k_range)),
                        more = vector("list", length(k_range)),
                        less = vector("list", length(k_range)))

for (k in k_range) {
  cat("    k =", k, "\n")
  ruv      <- RUVr(seqUQ, rownames(set), k = k, res)
  norm_cts <- normCounts(ruv)

  # RLE plot
  suppressMessages({
    pdf(paste0(outdir, "heart_RUVr_k", k, "_RLE.pdf"))
    plotRLE(ruv, outline = FALSE, ylim = c(-2, 2),
            col = colors, main = paste0("Heart snATAC RUVr k=", k))
    dev.off()
  })

  # Post-batch PCA
  log_post <- log1p(t(norm_cts))
  pca_k    <- prcomp(log_post, scale. = FALSE)
  pct_k    <- round(summary(pca_k)$importance[2, 1:2] * 100, 1)
  df_k     <- data.frame(
    PC1    = pca_k$x[, 1],
    PC2    = pca_k$x[, 2],
    sex    = meta[rownames(pca_k$x), "sex"],
    sample = rownames(pca_k$x)
  )
  p_k <- ggplot(df_k, aes(PC1, PC2, color = sex, label = sample)) +
    geom_point(size = 3) +
    geom_text_repel(size = 2.5, max.overlaps = 20) +
    scale_color_manual(values = c(female = "#E64B35", male = "#4DBBD5")) +
    labs(title = paste0("Heart snATAC RUVr k=", k),
         x = paste0("PC1 (", pct_k[1], "%)"),
         y = paste0("PC2 (", pct_k[2], "%)")) +
    theme_bw() +
    theme(plot.title = element_text(hjust = 0.5, face = "bold"))
  ggsave(paste0(outdir, "heart_RUVr_k", k, "_PCA.pdf"), p_k, width = 8, height = 6)

  # DAR: exactTest (female vs male)
  dgList_ruv <- DGEList(counts = norm_cts, genes = rownames(norm_cts), group = x)
  design_mat <- model.matrix(~ 0 + dge$samples$group)
  colnames(design_mat) <- levels(dge$samples$group)
  dgList_ruv <- estimateCommonDisp(dgList_ruv, design = design_mat)
  dgList_ruv <- estimateTagwiseDisp(dgList_ruv)

  et   <- exactTest(dgList_ruv, pair = c("male", "female"))
  topt <- topTags(et, n = nrow(et))

  cpm_ruv           <- cpm(norm_cts)
  result            <- cbind(as.data.frame(topt$table),
                             cpm_ruv[rownames(topt$table), ])
  result$sig        <- "NS"
  result$sig[result$FDR < fdr_cutoff & result$logFC >  fc_cutoff] <- "MORE"
  result$sig[result$FDR < fdr_cutoff & result$logFC < -fc_cutoff] <- "LESS"

  result_f          <- result[result$sig != "NS", ]
  result_more       <- result_f[result_f$sig == "MORE", ]
  result_less       <- result_f[result_f$sig == "LESS", ]

  cat("      MORE (female > male):", nrow(result_more), "\n")
  cat("      LESS (female < male):", nrow(result_less), "\n")

  write.csv(result,       paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, ".csv"))
  write.csv(result_f,     paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_filter.csv"))
  write.csv(result_more,  paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_MORE.csv"))
  write.csv(result_less,  paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_LESS.csv"))
  write.csv(cpm_ruv,      paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_normCPM.csv"))

  ki <- which(k_range == k)
  dar_lists$all[[ki]]  <- rownames(result_f)
  dar_lists$more[[ki]] <- rownames(result_more)
  dar_lists$less[[ki]] <- rownames(result_less)

  results_summary <- rbind(results_summary, data.frame(
    k           = k,
    total_peaks = nrow(result),
    MORE        = nrow(result_more),
    LESS        = nrow(result_less),
    total_DAR   = nrow(result_f)
  ))
}

write.csv(results_summary,
  paste0(outdir, "heart_DAR_female_vs_male_k_comparison_summary.csv"),
  row.names = FALSE)
cat("[6] k summary saved.\n")

# ============================================================
# 7. VOLCANO PLOT for each k
# ============================================================
cat("[7] Plotting volcano plots for each k...\n")

for (k in k_range) {
  result <- read.csv(
    paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, ".csv"),
    row.names = 1)

  p_vol <- ggplot(result, aes(logFC, -log10(FDR), color = sig)) +
    geom_point(size = 1, alpha = 0.6) +
    geom_vline(xintercept = c(-fc_cutoff, fc_cutoff), linetype = "dashed") +
    geom_hline(yintercept = -log10(fdr_cutoff), linetype = "dashed") +
    scale_color_manual(values = c(MORE = "#E64B35", LESS = "#4DBBD5", NS = "grey70")) +
    labs(title = paste0("Heart snATAC | female vs male (RUVr k=", k, ")"),
         x = "logFC (female - male)", y = "-log10(FDR)") +
    theme_bw() +
    theme(plot.title = element_text(hjust = 0.5, face = "bold"))
  ggsave(paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_volcano.pdf"),
         p_vol, width = 7, height = 6)
}

# ============================================================
# 8. DAR OVERLAP ACROSS k VALUES
# ============================================================
cat("[8] Computing DAR overlap across k values...\n")
k_labels <- paste0("k", k_range)
n_k      <- length(k_range)

make_overlap_mat <- function(lst, labels) {
  mat <- matrix(0L, n_k, n_k, dimnames = list(labels, labels))
  for (i in 1:n_k) for (j in 1:n_k)
    mat[i, j] <- length(intersect(lst[[i]], lst[[j]]))
  mat
}

ov_all  <- make_overlap_mat(dar_lists$all,  k_labels)
ov_more <- make_overlap_mat(dar_lists$more, k_labels)
ov_less <- make_overlap_mat(dar_lists$less, k_labels)

write.csv(ov_all,  paste0(outdir, "heart_DAR_overlap_all_k.csv"))
write.csv(ov_more, paste0(outdir, "heart_DAR_overlap_MORE_k.csv"))
write.csv(ov_less, paste0(outdir, "heart_DAR_overlap_LESS_k.csv"))

all_peaks <- unique(unlist(dar_lists$all))
peak_tbl  <- as.data.frame(
  matrix(0L, length(all_peaks), n_k, dimnames = list(all_peaks, k_labels)))
for (i in 1:n_k) if (length(dar_lists$all[[i]]) > 0)
  peak_tbl[dar_lists$all[[i]], i] <- 1L
peak_tbl$n_k       <- rowSums(peak_tbl[, k_labels])
peak_tbl$direction <- sapply(all_peaks, function(p) {
  is_more <- any(sapply(dar_lists$more, function(v) p %in% v))
  is_less <- any(sapply(dar_lists$less, function(v) p %in% v))
  if (is_more & is_less) "mixed" else if (is_more) "MORE" else "LESS"
})
peak_tbl <- peak_tbl[order(-peak_tbl$n_k), ]
write.csv(peak_tbl, paste0(outdir, "heart_DAR_peak_presence_across_k.csv"))

cat("[Done]\n")
