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

# ============================================================
# 0. CONFIG
# ============================================================
outdir  <- "/BRC/yan/heart/analysis/sn_ATAC_analysis/"
infile  <- paste0(outdir, "count_matrix_snPeak.txt")
k_range <- 1:6   # RUVr k values to test
k_final <- 3     # final k after checking RLE/PCA plots

# ============================================================
# 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"))
cat("    female:", sum(meta$sex == "female"),
    " male:", sum(meta$sex == "male"), "\n")

# ============================================================
# 3. FILTER
# ============================================================
cat("[3] Filtering low-count peaks...\n")
y    <- DGEList(counts = as.matrix(counts), group = x)
keep <- filterByExpr(y, group = x)
y    <- y[keep, , keep.lib.sizes = FALSE]
cat("    Peaks after filtering:", nrow(y), "\n")

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

y_tmp  <- calcNormFactors(y, method = "RLE")
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")

# ============================================================
# 5. RUVr k=1~6: diagnostic plots + save corrected matrix
# ============================================================
cat("[5] Running RUVr k=1~6...\n")
colors <- brewer.pal(8, "Set2")[as.integer(x)]

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

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

  # PCA (RUVSeq built-in)
  suppressMessages({
    pdf(paste0(outdir, "heart_RUVr_k", k, "_PCA.pdf"))
    plotPCA(set2, col = colors,
            main = paste0("Heart RUVr k=", k), cex = 0.6)
    dev.off()
  })

  # ggplot PCA with sample labels
  log_mat <- log1p(t(ndddd))
  pca     <- prcomp(log_mat, scale. = TRUE)
  pct     <- round(summary(pca)$importance[2, 1:2] * 100, 1)
  df_pca  <- data.frame(
    PC1    = pca$x[, 1],
    PC2    = pca$x[, 2],
    sex    = meta$sex,
    sample = meta$sample
  )
  p <- ggplot(df_pca, 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 RUVr k=", k),
         x = paste0("PC1 (", pct[1], "%)"),
         y = paste0("PC2 (", pct[2], "%)")) +
    theme_bw() +
    theme(plot.title = element_text(hjust = 0.5, face = "bold"))
  ggsave(paste0(outdir, "heart_RUVr_k", k, "_PCA_gg.pdf"),
         p, width = 8, height = 6)

  # save corrected matrix
  write.table(ndddd,
              paste0(outdir, "heart_RUVr_k", k, ".bed"),
              sep = "\t", quote = FALSE)
}

# ============================================================
# 6. DAR with chosen k
# ============================================================
cat("[6] Running DAR with k =", k_final, "...\n")
corrected <- read.table(paste0(outdir, "heart_RUVr_k", k_final, ".bed"),
                        header = TRUE, sep = "\t")
countdata <- round(corrected * 30)

y2    <- DGEList(counts = as.matrix(countdata), group = x)
keep2 <- filterByExpr(y2, group = x)
y2    <- y2[keep2, , keep.lib.sizes = FALSE]
cat("    Peaks after filtering:", nrow(y2), "\n")

y2 <- estimateCommonDisp(y2, verbose = FALSE)
y2 <- estimateTagwiseDisp(y2)

et  <- exactTest(y2)
dar <- topTags(et, n = Inf, adjust.method = "BH")$table
dar$peak      <- rownames(dar)
dar$direction <- "non-DAR"
dar$direction[dar$FDR < 0.01 & dar$logFC >  log2(1.5)] <- "MORE"
dar$direction[dar$FDR < 0.01 & dar$logFC < -log2(1.5)] <- "LESS"

cat("    MORE DAR:", sum(dar$direction == "MORE"), "\n")
cat("    LESS DAR:", sum(dar$direction == "LESS"), "\n")

# ============================================================
# 7. SAVE DAR results
# ============================================================
cat("[7] Saving results...\n")
write.table(dar,
  paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k_final, ".txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

write.table(dar[dar$direction != "non-DAR", ],
  paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k_final, "_filter.txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

cat("[Done]\n")
