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

# ============================================================
# 0. CONFIG
# ============================================================
outdir   <- "/home/bemiao/TaRGET_ATAC/All/analysis/"
infile   <- "/home/bemiao/TaRGET_ATAC/All/table/COUNT_MATRIX.txt"   # raw count matrix
smpfile  <- "/home/bemiao/TaRGET_ATAC/All/table/SAMPLE_TABLE.txt"   # sample table (V1=prefix, V2=condition, V3=n_rep)
peakdir  <- "/home/bemiao/TaRGET_ATAC/All/Peak/Lab_Peak/adult_liver/"

labs     <- c("BA", "BI", "AL", "DO", "WK", "MU", "ZB")
k_range  <- 1:4
k_final  <- 3

fdr_cutoff <- 0.01
fc_cutoff  <- log2(1.5)

# DAR comparisons: each row = c(lab, cond1_key1, cond1_key2, cond2_key1, cond2_key2)
# 仿照 script2 的 args[1:5]，在这里统一列出所有要跑的比较
comparisons <- list(
  c("M", "BPA10", "Li", "VEH", "Li"),
  c("M", "BPAlow", "Li", "VEH", "Li"),
  c("F", "BPA10", "Li", "VEH", "Li"),
  c("F", "BPAlow", "Li", "VEH", "Li")
)

# ============================================================
# 1. LOAD DATA
# ============================================================
cat("[1] Loading count matrix...\n")
peak <- read.table(infile, header = TRUE, sep = "\t")
cat("    Peaks:", nrow(peak), " Samples:", ncol(peak), "\n")

table_all <- read.table(smpfile)

# ============================================================
# 2. PER-LAB RUVr (k=1~4) + save normCounts
# ============================================================
cat("[2] Running per-lab RUVr...\n")

corrected_list <- list()   # 存每个 lab k_final 的 normCounts，后面合并

for (lab_name in labs) {
  cat("  [Lab]", lab_name, "\n")

  # --- 2a. 提取该 lab 的样本 ---
  lab_table <- table_all[grep(lab_name, table_all$V1), ]
  condition <- c()
  for (i in seq_len(nrow(lab_table))) {
    hits  <- grep(lab_table[i, 1], colnames(peak))
    names <- colnames(peak)[hits]
    sub   <- grep(lab_table[i, 2], names)
    condition <- c(condition, names[sub])
  }
  countdata <- peak[, condition]

  cons <- as.character(paste(lab_table$V1, lab_table$V2, sep = "_"))
  x    <- as.factor(c(rep(cons, lab_table$V3)))

  colors_lab <- brewer.pal(8, "Set2")[as.integer(x)]

  # --- 2b. 初始 edgeR：归一化 + 估计离散度 ---
  set <- newSeqExpressionSet(as.matrix(countdata),
           phenoData = data.frame(x, row.names = colnames(countdata)))
  design <- model.matrix(~ x, data = pData(set))

  y <- DGEList(counts = as.matrix(countdata), group = x)
  y <- calcNormFactors(y, method = "RLE")
  y <- estimateGLMCommonDisp(y, design)
  y <- estimateGLMTagwiseDisp(y, design)

  # --- 2c. 计算残差 ---
  set   <- newSeqExpressionSet(as.matrix(y$counts),
             phenoData = data.frame(x, row.names = colnames(countdata)))
  fit   <- glmFit(y, design)
  res   <- residuals(fit, type = "deviance")
  seqUQ <- betweenLaneNormalization(set, which = "upper")

  # --- 2d. raw QC plots ---
  pdf(paste0(outdir, lab_name, "_raw_RLE.pdf"))
  plotRLE(set, outline = FALSE, ylim = c(-2, 2),
          col = colors_lab, main = paste0(lab_name, " raw"))
  dev.off()

  pdf(paste0(outdir, lab_name, "_raw_PCA.pdf"))
  plotPCA(set, col = colors_lab,
          main = paste0(lab_name, " raw"), cex = 0.4)
  dev.off()

  write.table(y$counts,
    paste0(outdir, lab_name, "_raw.bed"),
    sep = "\t", quote = FALSE)

  # --- 2e. RUVr k=1~4 ---
  for (k in k_range) {
    set2  <- RUVr(seqUQ, rownames(set), k = k, res)
    ndddd <- normCounts(set2)

    # RUVSeq QC plots
    pdf(paste0(outdir, lab_name, "_RUVr_k", k, "_RLE.pdf"))
    plotRLE(set2, outline = FALSE, ylim = c(-2, 2),
            col = colors_lab, main = paste0(lab_name, " RUVr k=", k))
    dev.off()

    pdf(paste0(outdir, lab_name, "_RUVr_k", k, "_PCA.pdf"))
    plotPCA(set2, col = colors_lab,
            main = paste0(lab_name, " RUVr k=", k), cex = 0.4)
    dev.off()

    # ggplot2 PCA
    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],
      group  = x[match(rownames(pca$x), colnames(ndddd))],
      sample = rownames(pca$x)
    )
    p <- ggplot(df_pca, aes(PC1, PC2, color = group, label = sample)) +
      geom_point(size = 3) +
      geom_text_repel(size = 2.5, max.overlaps = 20) +
      labs(title = paste0(lab_name, " 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, lab_name, "_RUVr_k", k, "_PCA_gg.pdf"),
           p, width = 8, height = 6)

    write.table(ndddd,
      paste0(outdir, lab_name, "_RUVr_k", k, "_normCounts.bed"),
      sep = "\t", quote = FALSE, col.names = NA)

    # 保存 k_final 供后续合并
    if (k == k_final) {
      corrected_list[[lab_name]] <- ndddd
    }
  }
}

# ============================================================
# 3. MERGE corrected matrices across labs
# ============================================================
cat("[3] Merging corrected counts across labs...\n")

all_peaks <- Reduce(intersect, lapply(corrected_list, rownames))
cat("    Common peaks:", length(all_peaks), "\n")

merged <- do.call(cbind, lapply(corrected_list, function(m) m[all_peaks, ]))
cat("    Total samples after merge:", ncol(merged), "\n")

write.table(merged,
  paste0(outdir, "LAB_liver_correct.txt"),
  sep = "\t", quote = FALSE, col.names = NA)

# ============================================================
# 4. DAR (per comparison)
# ============================================================
cat("[4] Running DAR...\n")

for (cmp in comparisons) {
  lab_prefix <- cmp[1]
  cond1_k1   <- cmp[2];  cond1_k2 <- cmp[3]
  cond2_k1   <- cmp[4];  cond2_k2 <- cmp[5]
  con1 <- paste0(cond1_k1, "_", cond1_k2)
  con2 <- paste0(cond2_k1, "_", cond2_k2)
  cat("  ", lab_prefix, "|", con1, "vs", con2, "\n")

  # --- 4a. 读 lab-specific peak 白名单 ---
  peak_file <- paste0(peakdir, lab_prefix, "_Li_adt_peak.txt")
  lp        <- as.character(read.table(peak_file)$V1)
  pk        <- rownames(merged)
  ll        <- lp[lp %in% pk]
  mat       <- merged[ll, ]

  # --- 4b. 提取样本列 ---
  hits1  <- grep(cond1_k1, colnames(mat))
  names1 <- colnames(mat)[hits1]
  cond1_name <- names1[grep(cond1_k2, names1)]
  cond1_name <- cond1_name[grep(lab_prefix, cond1_name)]

  hits2  <- grep(cond2_k1, colnames(mat))
  names2 <- colnames(mat)[hits2]
  cond2_name <- names2[grep(cond2_k2, names2)]

  cn        <- c(cond1_name, cond2_name)
  countdata <- mat[, cn]
  countdata <- round(countdata * 30)   # 还原为整数 count

  n_cond1 <- length(cond1_name)
  n_cond2 <- length(cond2_name)

  # --- 4c. 低表达过滤 ---
  keep      <- rowSums(cpm(countdata) > 2) >= n_cond1
  countdata <- countdata[keep, ]

  # --- 4d. DGEList + exactTest ---
  x_dar <- factor(
    c(rep(con1, n_cond1), rep(con2, n_cond2)),
    levels = c(con2, con1)
  )
  wtest <- DGEList(counts = countdata, group = x_dar,
                   genes = rownames(countdata))
  wtest <- estimateCommonDisp(wtest, verbose = FALSE)
  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)), ]
  wftab <- cbind(wstop, wsexp)
  wftab <- wftab[order(wftab$PValue), ]

  # 修复重复列名
  colnames(wftab) <- c(
    "genes", "logFC", "logCPM", "PValue", "FDR",
    paste0(con1, "_", seq_len(n_cond1)),
    paste0(con2, "_", seq_len(n_cond2))
  )

  # --- 4e. ggplot2 PCA ---
  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],
    group  = x_dar[match(rownames(pca$x), colnames(countdata))],
    sample = rownames(pca$x)
  )
  p <- ggplot(df_pca, aes(PC1, PC2, color = group, label = sample)) +
    geom_point(size = 3) +
    geom_text_repel(size = 2.5, max.overlaps = 20) +
    labs(title = paste0(lab_prefix, " | ", con1, " vs ", con2),
         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, lab_prefix, "_DAR_", con1, "_", con2,
                "_RUVr_k", k_final, "_PCA.pdf"),
         p, width = 8, height = 6)

  # --- 4f. volcano plot ---
  wftab$label <- ifelse(
    abs(wftab$logFC) > fc_cutoff & wftab$FDR < fdr_cutoff,
    wftab$genes, ""
  )
  wftab$sig <- ifelse(
    abs(wftab$logFC) > fc_cutoff & wftab$FDR < fdr_cutoff,
    ifelse(wftab$logFC > 0, "MORE", "LESS"), "NS"
  )
  p2 <- 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(lab_prefix, " | ", con1, " vs ", con2),
         x = "logFC", y = "-log10(FDR)") +
    theme_bw() +
    theme(plot.title = element_text(hjust = 0.5, face = "bold"))
  ggsave(paste0(outdir, lab_prefix, "_DAR_", con1, "_", con2,
                "_RUVr_k", k_final, "_volcano.pdf"),
         p2, width = 7, height = 6)

  # --- 4g. 保存结果 ---
  prefix <- paste0(outdir, lab_prefix, "_DAR_", con1, "_", con2,
                   "_RUVr_k", k_final)
  write.table(wftab,
    paste0(prefix, ".txt"),
    sep = "\t", quote = FALSE, row.names = FALSE)

  wftab_sig <- wftab[abs(wftab$logFC) > fc_cutoff & wftab$FDR < fdr_cutoff, ]
  write.table(wftab_sig,
    paste0(prefix, "_filter.txt"),
    sep = "\t", quote = FALSE, row.names = FALSE)

  cat("    MORE:", sum(wftab$sig == "MORE", na.rm = TRUE),
      " LESS:", sum(wftab$sig == "LESS", na.rm = TRUE), "\n")
}

cat("[Done]\n")