#!/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   # RUVr k values to test
k_final <- 6     # 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  <- 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")

# Bug5 fix: use sex-based colors instead of brewer index
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)

  # 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
  # Bug1 fix: scale. = FALSE to match RUVSeq plotPCA behavior
  log_mat <- log1p(t(ndddd))
  pca     <- prcomp(log_mat, scale. = FALSE)
  pct     <- round(summary(pca)$importance[2, 1:2] * 100, 1)

  # Bug2 fix: align meta by rownames of pca$x
  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 with RUV batch correction (k_final, GLM LRT)
# ============================================================
cat("[6] Running DAR with RUV batch correction (k =", k_final, ")...\n")

fdr_cutoff <- 0.01
fc_cutoff  <- log2(1.5)

n_female <- sum(meta$sex == "female")
n_male   <- sum(meta$sex == "male")

# W factors from RUVr at k_final, used as covariates in GLM
set_final   <- RUVr(seqUQ, rownames(set), k = k_final, res)
W           <- pData(set_final)[, grep("W_", colnames(pData(set_final))), drop = FALSE]

group_edgeR <- factor(meta$sex, levels = c("male", "female"))
design_ruv  <- model.matrix(~ group_edgeR + W)

# Create DGEList (peaks already filtered in step 3)
wtest <- DGEList(counts = y$counts,
                 group  = group_edgeR,
                 genes  = rownames(y$counts))
wtest <- calcNormFactors(wtest, method = "RLE")
wtest <- estimateDisp(wtest, design_ruv)
fit   <- glmFit(wtest, design_ruv)
lrt   <- glmLRT(fit, coef = "group_edgeRfemale")

# CPM for output table
wcpm  <- cpm(wtest)

wntop  <- topTags(lrt, 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", "LR", "PValue", "FDR",
                     rep("female", n_female), rep("male", n_male))

wftab_selected <- wftab[abs(wftab$logFC) > fc_cutoff & wftab$FDR < fdr_cutoff, ]

cat("    MORE DAR (female>male):", sum(wftab_selected$logFC > 0), "\n")
cat("    LESS DAR (female<male):", sum(wftab_selected$logFC < 0), "\n")

# ============================================================
# 7. SAVE DAR results
# ============================================================
cat("[7] Saving results...\n")
write.table(wftab,
  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")
