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

# ============================================================
# 0. CONFIG
# ============================================================
outdir  <- "/BRC/yan/heart/analysis/sn_ATAC_analysis/results_v5/"
infile  <- paste0(outdir, "count_matrix_snPeak.txt")
k_range <- 1:6

cpm_cutoff <- 1
fdr_cutoff <- 0.05
fc_cutoff  <- log2(1.5)   # 0.58

# ============================================================
# 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")

colors <- ifelse(meta$sex == "female", "#E64B35", "#4DBBD5")

# ============================================================
# 3. FILTER — CPM > 1 in at least min(group size) samples
# ============================================================
cat("[3] Filtering low-count peaks (CPM >", cpm_cutoff, ")...\n")
dge   <- DGEList(counts = as.matrix(counts), group = x)
cpm_mat <- cpm(counts)
keep    <- rowSums(cpm_mat > cpm_cutoff) >= min(table(x))
dge     <- dge[keep, , keep.lib.sizes = FALSE]
cpm_mat <- cpm_mat[keep, ]
cat("    Peaks after filtering:", nrow(dge), "\n")

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

y_tmp  <- DGEList(counts = counts(set), group = x)
y_tmp  <- calcNormFactors(y_tmp, method = "upperquartile")
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: diagnostics + full DAR per k
# ============================================================
cat("[5] Running RUVr k=1~6 with full DAR each...\n")

results_summary <- data.frame()

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

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

  # --- ggplot PCA ---
  log_mat <- log1p(t(norm_cts))
  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 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.pdf"), p_pca, width = 8, height = 6)

  # --- DAR: exactTest (female vs male) ---
  dgList_ruv <- DGEList(counts = norm_cts,
                        genes  = rownames(norm_cts),
                        group  = x)

  design_mat <- model.matrix(~ 0 + dge$samples$group)
  colnames(design_mat) <- levels(dge$samples$group)

  dgList_ruv <- estimateCommonDisp(dgList_ruv, design = design_mat)
  dgList_ruv <- estimateTagwiseDisp(dgList_ruv)

  et   <- exactTest(dgList_ruv, pair = c("male", "female"))
  topt <- topTags(et, n = nrow(et))

  cpm_ruv <- cpm(norm_cts)
  cpm_ruv <- cpm_ruv[rownames(topt$table), ]

  result           <- cbind(as.data.frame(topt$table), cpm_ruv)
  result$sig       <- ifelse(
    abs(result$logFC) > fc_cutoff & result$FDR < fdr_cutoff,
    ifelse(result$logFC > 0, "MORE", "LESS"), "NS"
  )

  result_fdr001    <- result[result$FDR < 0.001, ]
  result_fdr05_fc  <- result[result$FDR < fdr_cutoff & abs(result$logFC) > fc_cutoff, ]

  mean_var <- mean(apply(norm_cts, 1, var))
  cat("      Mean variance:", round(mean_var, 2),
      " | FDR<0.001:", nrow(result_fdr001),
      " | FDR<0.05+FC:", nrow(result_fdr05_fc), "\n")

  write.csv(result,
    paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_exactTest.csv"))
  write.csv(result_fdr001,
    paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_exactTest_FDR001.csv"))
  write.csv(result_fdr05_fc,
    paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_exactTest_FDR05_FC.csv"))

  results_summary <- rbind(results_summary, data.frame(
    k             = k,
    total_peaks   = nrow(result),
    FDR001        = nrow(result_fdr001),
    FDR05_FC      = nrow(result_fdr05_fc),
    MORE          = sum(result$sig == "MORE"),
    LESS          = sum(result$sig == "LESS"),
    mean_variance = mean_var
  ))
}

write.csv(results_summary,
  paste0(outdir, "heart_DAR_female_vs_male_k_comparison_summary.csv"),
  row.names = FALSE)
cat("[5] k summary saved.\n")

# ============================================================
# 6. VOLCANO PLOT for each k (FDR < 0.05, |logFC| > 0.58)
# ============================================================
cat("[6] Plotting volcano plots for each k...\n")

for (k in k_range) {
  result <- read.csv(
    paste0(outdir, "heart_DAR_female_vs_male_RUVr_k", k, "_exactTest.csv"),
    row.names = 1)

  p_vol <- ggplot(result, 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, ")"),
         x = "logFC (female - male)", 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, "_volcano.pdf"),
         p_vol, width = 7, height = 6)
}

cat("[Done]\n")
