#!/usr/bin/env Rscript

# =====================================================
# JITTER PLOTS:
# TE CONSERVATION vs
# DNA METHYLATION / EXPRESSION
# =====================================================

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

setDTthreads(8)

# =====================================================
# INPUT
# =====================================================

infile <- "/BLUES/eric/ONT_WGBS/Figure_3/Jitter_plot_conservation/TE_conservation_methylation_expression.tsv"

outdir <- "/BLUES/eric/ONT_WGBS/Figure_3/Jitter_plot_conservation"

# =====================================================
# STATE
# =====================================================

# Choose:
# "Primed"
# "Naive"
# "TSC"

state_name <- "Primed"

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

dt <- fread(infile)

cat("Loaded rows:", nrow(dt), "\n")

# =====================================================
# DYNAMIC COLUMN NAMES
# =====================================================

meth_col <- paste0(state_name, "_methylation")
expr_col <- paste0(state_name, "_expression")

cat("Using methylation column:", meth_col, "\n")
cat("Using expression column:", expr_col, "\n")

# =====================================================
# COLORS
# =====================================================

group_colors <- c(
  "Amniota"  = "#F8766D",
  "Mammalia" = "#00BFC4",
  "Primate"  = "#C77CFF",
  "Homo"     = "#7CAE00"
)

# =====================================================
# SHAPES
# =====================================================

shape_vals <- c(
  "DNA"  = 17,
  "LINE" = 15,
  "LTR"  = 16,
  "SINE" = 18
)

# =====================================================
# LABELS
# =====================================================

# -----------------------------
# Upper panel candidates
# -----------------------------

label_meth <- dt[
  Group == "Homo" |
    (
      Group %in% c("Mammalia", "Primate") &
        get(meth_col) < 0.75
    )
]

# -----------------------------
# Lower panel candidates
# -----------------------------

expr_cut_mammalia <- quantile(
  dt[Group == "Mammalia"][[expr_col]],
  0.985,
  na.rm = TRUE
)

expr_cut_primate <- quantile(
  dt[Group == "Primate"][[expr_col]],
  0.995,
  na.rm = TRUE
)

label_expr <- dt[
  Group == "Homo" |
    (
      Group == "Mammalia" &
        get(expr_col) >= expr_cut_mammalia
    ) |
    (
      Group == "Primate" &
        get(expr_col) >= expr_cut_primate
    )
]

# =====================================================
# COMBINE LABELS FROM BOTH PANELS
# =====================================================

all_labels <- unique(
  rbind(
    label_meth,
    label_expr,
    fill = TRUE
  ),
  by = "subfamily"
)

cat("Total combined labels:",
    nrow(all_labels), "\n")

# =====================================================
# THEME
# =====================================================

theme_te <- theme_classic(base_size = 14) +
  theme(
    plot.title = element_text(
      face = "bold",
      hjust = 0.5,
      size = 16
    ),

    axis.title = element_text(
      face = "bold",
      size = 14
    ),

    axis.text = element_text(
      size = 12
    ),

    legend.title = element_text(
      face = "bold",
      size = 13
    ),

    legend.text = element_text(
      size = 11
    )
  )

# =====================================================
# GENERATE ALL STATES
# =====================================================

for (state_name in c("Primed", "Naive", "TSC")) {

  cat("\n========================\n")
  cat("Processing:", state_name, "\n")
  cat("========================\n")

  meth_col <- paste0(state_name, "_methylation")
  expr_col <- paste0(state_name, "_expression")

  # --------------------------------
  # STATE-SPECIFIC OUTLIER LABELS
  # --------------------------------

  meth_vals <- dt[[meth_col]]
  expr_vals <- log10(dt[[expr_col]] + 1)

  # --------------------------------
  # METHYLATION OUTLIERS
  # --------------------------------

  meth_low_cut <- quantile(
    meth_vals,
    0.01,
    na.rm = TRUE
  )

  meth_high_cut <- quantile(
    meth_vals,
    0.99,
    na.rm = TRUE
  )

  # --------------------------------
  # EXPRESSION OUTLIERS
  # --------------------------------

  expr_high_cut <- quantile(
    expr_vals,
    0.995,
    na.rm = TRUE
  )

  # --------------------------------
  # UPPER PANEL LABELS
  # --------------------------------

  label_meth <- dt[
    Group == "Homo" |
      (
        Group %in% c("Mammalia", "Primate") &
          (
            get(meth_col) <= meth_low_cut |
            get(meth_col) >= meth_high_cut
          )
      )
  ]

  # --------------------------------
  # LOWER PANEL LABELS
  # --------------------------------

  label_expr <- dt[
    Group == "Homo" |
      (
        Group %in% c("Mammalia", "Primate") &
          log10(get(expr_col) + 1) >= expr_high_cut
      )
  ]

  # --------------------------------
  # UNION OF BOTH LABEL SETS
  # --------------------------------

  all_labels <- unique(
    rbind(
      label_meth,
      label_expr,
      fill = TRUE
    ),
    by = "subfamily"
  )

  cat("Labels for", state_name, ":",
      nrow(all_labels), "\n")

  # --------------------------------
  # METHYLATION
  # --------------------------------

  p_meth <- ggplot(
    dt,
    aes(
      x = (-1) * MYA,
      y = get(meth_col),
      color = Group,
      shape = class
    )
  ) +

    geom_jitter(
      width = 2,
      height = 0.01,
      alpha = 0.8,
      size = 2
    ) +

    geom_text_repel(
      data = all_labels,
      aes(label = subfamily),
      size = 3,
      max.overlaps = Inf,
      box.padding = 0.4,
      point.padding = 0.2,
      segment.size = 0.3,
      min.segment.length = 0
    ) +

    scale_color_manual(
      values = group_colors
    ) +

    scale_shape_manual(
      values = shape_vals
    ) +

    labs(
      title = paste0(
        "TE evolutionary age vs DNA methylation in ",
        state_name
      ),
      x = "Insertion time of TE (million years)",
      y = "DNA methylation"
    ) +

    theme_te

  # --------------------------------
  # EXPRESSION
  # --------------------------------

  p_expr <- ggplot(
    dt,
    aes(
      x = (-1) * MYA,
      y = log10(get(expr_col) + 1),
      color = Group,
      shape = class
    )
  ) +

    geom_jitter(
      width = 2,
      height = 0.05,
      alpha = 0.8,
      size = 2
    ) +

    geom_text_repel(
      data = all_labels,
      aes(label = subfamily),
      size = 3,
      max.overlaps = Inf,
      box.padding = 0.4,
      point.padding = 0.2,
      segment.size = 0.3,
      min.segment.length = 0
    ) +

    scale_color_manual(
      values = group_colors
    ) +

    scale_shape_manual(
      values = shape_vals
    ) +

    labs(
      title = paste0(
        "TE evolutionary age vs expression in ",
        state_name
      ),
      x = "Insertion time of TE (million years)",
      y = "log10(mean FPKM + 1)"
    ) +

    theme_te

  combined <- p_meth / p_expr

  outfile <- file.path(
    outdir,
    paste0(
      "TE_conservation_jitter_plots_",
      state_name,
      ".pdf"
    )
  )

  ggsave(
    outfile,
    combined,
    width = 10,
    height = 12
  )

  cat("Saved:", outfile, "\n")
}