#!/usr/bin/env Rscript

suppressPackageStartupMessages({
  library(data.table)
  library(ggplot2)
  library(ggrepel)
  library(patchwork)
})

setDTthreads(8)

# =====================================================
# INPUT / OUTPUT
# =====================================================

outdir <- "/BLUES/eric/ONT_WGBS/Figure_3/MA_Plot/deseq"

infile <- file.path(outdir, "RNAseq_log2FC_DESeq2.tsv")

pdf_out <- file.path(outdir, "TE_RNA_MA_DESeq2.pdf")
tsv_out <- file.path(outdir, "TE_RNA_MA_DESeq2_plot_table.tsv")

# =====================================================
# LOAD
# =====================================================

dt <- fread(infile)

# =====================================================
# PREPARE PLOT TABLE
# =====================================================

np <- dt[, .(
  subfamily,
  comparison = "Naive vs Primed",
  baseMean = baseMean_NP,
  mean_expression = log2(baseMean_NP + 1),
  log2FC = log2FC_NP,
  padj = padj_NP
)]

tn <- dt[, .(
  subfamily,
  comparison = "TSC vs Naive",
  baseMean = baseMean_TN,
  mean_expression = log2(baseMean_TN + 1),
  log2FC = log2FC_TN,
  padj = padj_TN
)]

plot_dt <- rbind(np, tn, fill = TRUE)

plot_dt <- plot_dt[
  is.finite(mean_expression) &
    is.finite(log2FC)
]

# =====================================================
# CUTOFFS
# =====================================================

fc_cut <- 2
padj_cut <- 0.05

plot_dt[, status := fifelse(
  log2FC >= fc_cut,
  "Up",
  fifelse(
    log2FC <= -fc_cut,
    "Down",
    "No change"
  )
)]

plot_dt[, significant := !is.na(padj) & padj < padj_cut]

fwrite(
  plot_dt,
  tsv_out,
  sep = "\t",
  quote = FALSE
)

# =====================================================
# PLOT FUNCTION
# =====================================================

plot_ma <- function(d, title_text, ylab_text) {

  d <- copy(d)

  d[, status := fifelse(
    log2FC >= fc_cut,
    "Up",
    fifelse(
      log2FC <= -fc_cut,
      "Down",
      "No change"
    )
  )]

  lab_up <- d[log2FC >= fc_cut]
  lab_down <- d[log2FC <= -fc_cut]

  ggplot(
    d,
    aes(
      x = mean_expression,
      y = log2FC
    )
  ) +
    geom_point(
      data = d[status == "No change"],
      color = "grey60",
      alpha = 0.6,
      size = 0.8
    ) +
    geom_point(
      data = lab_up,
      color = "#d62728",
      alpha = 0.9,
      size = 1.4
    ) +
    geom_point(
      data = lab_down,
      color = "#1f77b4",
      alpha = 0.9,
      size = 1.4
    ) +
    geom_hline(
      yintercept = 0,
      color = "black",
      linewidth = 0.4
    ) +
    geom_hline(
      yintercept = c(-fc_cut, fc_cut),
      color = "grey35",
      linetype = "dashed",
      linewidth = 0.5
    ) +
    geom_text_repel(
      data = lab_up,
      aes(label = subfamily),
      color = "#d62728",
      size = 1.8,
      max.overlaps = Inf,
      box.padding = 0.08,
      point.padding = 0.04,
      segment.size = 0.15,
      min.segment.length = 0.05,
      force = 0.5,
      force_pull = 1.5,
      max.time = 2
    ) +
    geom_text_repel(
      data = lab_down,
      aes(label = subfamily),
      color = "#1f77b4",
      size = 1.8,
      max.overlaps = Inf,
      box.padding = 0.08,
      point.padding = 0.04,
      segment.size = 0.15,
      min.segment.length = 0.05,
      force = 0.5,
      force_pull = 1.5,
      max.time = 2
    ) +
    annotate(
      "text",
      x = Inf,
      y = Inf,
      label = paste0("n = ", nrow(d)),
      hjust = 1.1,
      vjust = 1.5,
      size = 4.5,
      fontface = "bold"
    ) +
    labs(
      title = title_text,
      subtitle = paste0(
        "x-axis: log2(DESeq2 normalized mean counts + 1); ",
        "y-axis: DESeq2 log2 fold change; ",
        "dashed lines: |log2FC| = ", fc_cut
      ),
      x = "log2(normalized mean counts + 1)",
      y = ylab_text
    ) +
    theme(
      plot.title = element_text(
        face = "bold",
        hjust = 0.5,
        size = 16
      ),
      plot.subtitle = element_text(
        hjust = 0.5,
        size = 10
      ),
      axis.title = element_text(
        face = "bold",
        size = 14
      ),
      axis.text = element_text(size = 12)
    )
}

# =====================================================
# PLOT
# =====================================================

p1 <- plot_ma(
  plot_dt[comparison == "Naive vs Primed"],
  "RNA-seq DESeq2: Naive vs Primed",
  "DESeq2 log2FC (Naive - Primed)"
)

p2 <- plot_ma(
  plot_dt[comparison == "TSC vs Naive"],
  "RNA-seq DESeq2: TSC vs Naive",
  "DESeq2 log2FC (TSC - Naive)"
)

fig <- p1 / p2

ggsave(
  pdf_out,
  fig,
  width = 8,
  height = 10,
  useDingbats = FALSE
)

cat("Saved:\n")
cat(pdf_out, "\n")
cat(tsv_out, "\n")
