#!/usr/bin/env Rscript
###############################################################################
# DESeq2 差异表达分析模板
# 输入：run_rnaseq.sh 产出的 counts_matrix.tsv + samples.tsv
# 输出：标准化表达矩阵、差异结果表、火山图、PCA 图、样本距离热图
###############################################################################

suppressPackageStartupMessages({
  library(DESeq2)
  library(ggplot2)
  library(pheatmap)
})

## ------------------------------ 参数 ---------------------------------------
count_file   <- "results/04_counts/counts_matrix.tsv"   # 计数矩阵
coldata_file <- "samples.tsv"                            # 样本表（sample_id / group）
out_dir      <- "results/06_DE"
ref_level    <- "control"    # 对照组（参考水平），按你的分组改
padj_cut     <- 0.05         # 显著性阈值
lfc_cut      <- 1            # log2FC 阈值
dir.create(out_dir, showWarnings = FALSE, recursive = TRUE)

## ---------------------------- 读入数据 -------------------------------------
counts <- read.delim(count_file, row.names = 1, check.names = FALSE)
counts <- as.matrix(counts)

## 读样本表：表头写成 "sample_id\tgroup" 或 "# sample_id\tgroup" 都要能用。
## 别用 read.delim(comment.char = "#")——若表头是 "# sample_id ..."，整行会被当成注释吃掉，
## 列名整体错位，下一行 coldata$sample_id 直接变成 NULL 并报
## "replacement has 0 rows"，且报错信息完全看不出真正原因。
coldata_lines <- readLines(coldata_file, warn = FALSE)
coldata_lines <- coldata_lines[!grepl("^\\s*$", coldata_lines)]
coldata_lines[1] <- sub("^\\s*#\\s*", "", coldata_lines[1])       # 剥掉表头行可能带的 #
coldata_lines <- coldata_lines[!grepl("^\\s*#", coldata_lines)]   # 其余 # 行仍算注释
coldata <- read.delim(text = paste(coldata_lines, collapse = "\n"),
                      stringsAsFactors = FALSE)
rownames(coldata) <- coldata$sample_id
coldata$group <- factor(coldata$group)
coldata$group <- relevel(coldata$group, ref = ref_level)

## 对齐样本顺序（这一步漏了会直接报错或结果错乱）
common <- intersect(colnames(counts), rownames(coldata))
if (length(common) == 0) stop("计数矩阵列名与样本表 sample_id 对不上，请检查")
counts  <- counts[, common, drop = FALSE]
coldata <- coldata[common, , drop = FALSE]

message("样本数：", ncol(counts), "；基因数：", nrow(counts))
message("分组：", paste(levels(coldata$group), collapse = " vs "))

## --------------------------- 构建 DESeq 对象 -------------------------------
dds <- DESeqDataSetFromMatrix(countData = counts,
                              colData   = coldata,
                              design    = ~ group)

## 过滤低表达基因：至少在一半样本里 count >= 10
keep <- rowSums(counts(dds) >= 10) >= floor(ncol(counts) / 2)
dds  <- dds[keep, ]
message("过滤后保留基因数：", nrow(dds))

## ---------------------- 标准化 + 差异分析 ----------------------------------
dds <- DESeq(dds)

## 提取结果（本例取 实验组 vs 对照组）
res <- results(dds, contrast = c("group",
                                 setdiff(levels(coldata$group), ref_level)[1],
                                 ref_level))
res <- res[order(res$padj), ]

## 收缩 log2FC（推荐，减少低表达基因的假阳性）
res_shrunk <- lfcShrink(dds, coef = resultsNames(dds)[2], type = "apeglm")

## ------------------------------ 导出结果 -----------------------------------
norm_counts <- counts(dds, normalized = TRUE)
write.table(norm_counts, file = file.path(out_dir, "normalized_counts.tsv"),
            sep = "\t", quote = FALSE, col.names = NA)

res_df <- as.data.frame(res_shrunk)
res_df$gene <- rownames(res_df)
write.table(res_df, file = file.path(out_dir, "DE_results_all.tsv"),
            sep = "\t", quote = FALSE, row.names = FALSE)

## 显著差异基因
sig <- subset(res_df, !is.na(padj) & padj < padj_cut & abs(log2FoldChange) > lfc_cut)
write.table(sig, file = file.path(out_dir, "DE_genes_significant.tsv"),
            sep = "\t", quote = FALSE, row.names = FALSE)
message("显著差异基因：", nrow(sig), " 个（上调 ",
        sum(sig$log2FoldChange > 0), "，下调 ", sum(sig$log2FoldChange < 0), "）")

## ------------------------------ 可视化 -------------------------------------

## 1) 火山图
volcano_df <- as.data.frame(res_shrunk)
volcano_df$significance <- "NotSig"
volcano_df$significance[volcano_df$padj < padj_cut &
                        volcano_df$log2FoldChange >  lfc_cut] <- "Up"
volcano_df$significance[volcano_df$padj < padj_cut &
                        volcano_df$log2FoldChange < -lfc_cut] <- "Down"

p_volcano <- ggplot(volcano_df,
                    aes(x = log2FoldChange, y = -log10(padj),
                        color = significance)) +
  geom_point(alpha = 0.7, size = 1.2) +
  scale_color_manual(values = c(Up = "#c0392b", Down = "#2471a3",
                                NotSig = "grey75")) +
  geom_vline(xintercept = c(-lfc_cut, lfc_cut), linetype = "dashed") +
  geom_hline(yintercept = -log10(padj_cut), linetype = "dashed") +
  labs(x = "log2 fold change", y = "-log10(padj)", color = NULL,
       title = "Differential expression") +
  theme_bw(base_size = 13)

ggsave(file.path(out_dir, "volcano.pdf"), p_volcano, width = 6, height = 5)
ggsave(file.path(out_dir, "volcano.png"), p_volcano, width = 6, height = 5,
       dpi = 300)

## 2) PCA 图（用方差稳定变换后的数据）
## vst() 默认拿 1000 个基因(参数 nsub)估计离散度，过滤后基因数不足 1000 时会直接
## 报 "less than 'nsub' rows"，小规模测试数据（如 chr22 子集）必然踩到；
## 基因数不够时改用 varianceStabilizingTransformation，它用全部基因、不要求 nsub。
if (nrow(dds) >= 1000) {
  vsd <- vst(dds, blind = FALSE)
} else {
  vsd <- varianceStabilizingTransformation(dds, blind = FALSE)
}
pca <- plotPCA(vsd, intgroup = "group", returnData = TRUE)
pct <- round(100 * attr(pca, "percentVar"))

p_pca <- ggplot(pca, aes(PC1, PC2, color = group, label = name)) +
  geom_point(size = 3) +
  geom_text(vjust = -1, size = 3, show.legend = FALSE) +
  xlab(paste0("PC1: ", pct[1], "% variance")) +
  ylab(paste0("PC2: ", pct[2], "% variance")) +
  theme_bw(base_size = 13)

ggsave(file.path(out_dir, "PCA.pdf"), p_pca, width = 6, height = 5)

## 3) 样本相关性热图
sample_dist <- dist(t(assay(vsd)))
dist_mat <- as.matrix(sample_dist)
pdf(file.path(out_dir, "sample_distance_heatmap.pdf"), width = 7, height = 6)
pheatmap(dist_mat, clustering_distance_rows = sample_dist,
         clustering_distance_cols = sample_dist,
         annotation_col = data.frame(group = coldata$group,
                                     row.names = rownames(coldata)),
         main = "Sample-to-sample distance")
dev.off()

## 4) 显著基因热图（取前 50 个最显著基因）
if (nrow(sig) > 1) {
  top_n <- head(sig$gene, 50)
  pdf(file.path(out_dir, "top50_heatmap.pdf"), width = 8, height = 10)
  pheatmap(assay(vsd)[top_n, ], scale = "row",
           annotation_col = data.frame(group = coldata$group,
                                       row.names = rownames(coldata)),
           show_rownames = TRUE, fontsize_row = 6,
           main = "Top 50 DE genes")
  dev.off()
}

message("全部完成，结果见：", normalizePath(out_dir))

###############################################################################
# 依赖安装：
#   BiocManager::install(c("DESeq2"))
#   install.packages(c("ggplot2","pheatmap"))
# 运行：
#   Rscript deseq2_de.R
# 注意：
#   1) apeglm 收缩需要 install.packages("apeglm")；没装可把 type 改成 "ashr"
#      或直接不收缩，用 res（注意此时 log2FC 更保守）
#   2) 分组水平名必须是合法 R 变量名，避免 "ctrl-x" 这类带横线的名字
###############################################################################
