#!/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
k_final <- 6

fdr_cutoff <- 0.01
fc_cutoff  <- 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"))
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  <- estimateDisp(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 <- ifelse(meta$sex == "female", "#E64B35", "#4DBBD5")

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

  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()
  })

  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()
  })

  log_mat <- log1p(t(ndddd))
  pca     <- prcomp(log_mat, scale. = FALSE)
  pct     <- round(summary(pca)$importance[2, 1:2] * 100, 1)
  df_pca  <- data.frame(
    PC1    = pca$x[, 1],
    PC2    = pca$x[, 2],
    sex    = meta[rownames(pca$x), "sex"],
    sample = meta[rownames(pca$x), "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)

  write.table(ndddd,
              paste0(outdir, "heart_RUVr_k", k, "_normCounts.txt"),
              sep = "\t", quote = FALSE, col.names = NA)
}

# ============================================================
# 6. DAR on RUVr k_final normCounts (exactTest)
# ============================================================
cat("[6] Running DAR on RUVr k=", k_final, " normCounts...\n")

set_final   <- RUVr(seqUQ, rownames(set), k = k_final, res)
ndddd_final <- normCounts(set_final)
countdata   <- round(ndddd_final)   # integer counts required by edgeR

group_edgeR <- factor(meta$sex, levels = c("male", "female"))
n_female    <- sum(meta$sex == "female")
n_male      <- sum(meta$sex == "male")

# DGEList (no CPM filter: already filtered in step 3)
wtest <- DGEList(counts = countdata,
                 group  = group_edgeR,
                 genes  = rownames(countdata))
wtest <- estimateCommonDisp(wtest, verbose = TRUE)
wtest <- estimateTagwiseDisp(wtest)

wcpm  <- cpm(wtest$pseudo.counts)
wet   <- exactTest(wtest)
wntop <- topTags(wet, n = nrow(wtest))

wstop  <- wntop$table[order(rownames(wntop$table)), ]
wsexp  <- wcpm[order(rownames(wcpm)), ]
wnftab <- cbind(wstop, wsexp)
wftab  <- wnftab[order(wnftab$PValue), ]
colnames(wftab) <- c(
  "genes", "logFC", "logCPM", "PValue", "FDR",
  colnames(countdata)
)

wftab$sig <- ifelse(
  abs(wftab$logFC) > fc_cutoff & wftab$FDR < fdr_cutoff,
  ifelse(wftab$logFC > 0, "MORE", "LESS"), "NS"
)

cat("    MORE DAR (female>male):", sum(wftab$sig == "MORE"), "\n")
cat("    LESS DAR (female<male):", sum(wftab$sig == "LESS"), "\n")

# ============================================================
# 7. DAR PCA
# ============================================================
cat("[7] Plotting DAR PCA...\n")

log_mat <- log1p(t(countdata))
pca     <- prcomp(log_mat, scale. = FALSE)
pct     <- round(summary(pca)$importance[2, 1:2] * 100, 1)
df_pca  <- data.frame(
  PC1    = pca$x[, 1],
  PC2    = pca$x[, 2],
  sex    = meta[rownames(pca$x), "sex"],
  sample = rownames(pca$x)
)
p_pca <- 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 DAR PCA (RUVr k=", k_final, ")"),
       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_DAR_female_vs_male_RUVr_k", k_final, "_PCA.pdf"),
       p_pca, width = 8, height = 6)

# ============================================================
# 8. Volcano plot
# ============================================================
cat("[8] Plotting volcano...\n")

p_vol <- ggplot(wftab, 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 | female vs male (RUVr k=", k_final, ")"),
       x = "logFC", 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_final, "_volcano.pdf"),
       p_vol, width = 7, height = 6)

# ============================================================
# 9. SAVE results
# ============================================================
cat("[9] Saving results...\n")

out_cols       <- setdiff(colnames(wftab), c("sig"))
wftab_selected <- wftab[wftab$sig != "NS", out_cols]

write.table(wftab[, out_cols],
  file = paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k_final, ".txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

write.table(wftab_selected,
  file = paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k_final, "_filter.txt"),
  sep = "\t", quote = FALSE, row.names = FALSE)

cat("[Done]\n")
