#!/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("female", "male"))
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 with exactTest (RUVr-corrected counts as input)
# ============================================================
cat("[6] Running DAR with exactTest on RUVr normCounts...\n")

# Step 1: 取 k_final 的 normCounts
set_final        <- RUVr(seqUQ, rownames(set), k = k_final, res)
corrected_counts <- normCounts(set_final)

# Step 2: group 按列名对应，不依赖列顺序          ← Bug2 fix
group_edgeR <- factor(
  meta[colnames(corrected_counts), "sex"],
  levels = c("female", "male")
)
n_female <- sum(group_edgeR == "female")
n_male   <- sum(group_edgeR == "male")

# Step 3: DGEList，不做 calcNormFactors          ← normCounts 已归一化
wtest <- DGEList(
  counts = corrected_counts,
  group  = group_edgeR,
  genes  = rownames(corrected_counts)
)

# Step 4: 估计离散度
wtest <- estimateCommonDisp(wtest, verbose = TRUE)
wtest <- estimateTagwiseDisp(wtest)

# Step 5: exactTest (female vs male)
wcpm  <- cpm(wtest$pseudo.counts)
wet   <- exactTest(wtest, pair = c("female", "male"))
wntop <- topTags(wet, n = nrow(wtest))

wstop  <- wntop$table[order(rownames(wntop$table)), ]
wsexp  <- wcpm[order(rownames(wcpm)), ]
wftab  <- cbind(wstop, wsexp)
wftab  <- wftab[order(wftab$PValue), ]

# Step 6: 修复重复列名                           ← Bug3 fix
colnames(wftab) <- c(
  "genes", "logFC", "logCPM", "PValue", "FDR",
  paste0("female_", seq_len(n_female)),
  paste0("male_",   seq_len(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")