#!/usr/bin/env Rscript
suppressPackageStartupMessages({
  library(data.table)
  library(ggplot2)
  library(VennDiagram)
  library(grid)
  library(patchwork)
})

setDTthreads(8)

BASE <- "/BLUES/eric/ONT_WGBS/Allele_analysis"
OUT  <- file.path(BASE, "summary_plots_d0.5_new")
dir.create(OUT, showWarnings=FALSE, recursive=TRUE)

cols <- c(
  Prime = "#4C78A8",
  Naive = "#F58518",
  TSC   = "#54A24B"
)

read_dt <- function(f) {
  dt <- fread(f)
  dt[, chr := gsub('"', '', chr)]
  dt[, id := paste(chr, win_start, win_end, sep=":")]
  dt
}

keep_auto <- function(dt) {
  dt[chr %in% paste0("chr", 1:22)]
}

# =========================
# Load data
# =========================

p_pos <- read_dt(file.path(BASE, "window_delta_prime/prime_windows_1000bp_delta_gt_0.5.tsv"))
p_neg <- read_dt(file.path(BASE, "window_delta_prime/prime_windows_1000bp_delta_lt_-0.5.tsv"))

n_pos <- read_dt(file.path(BASE, "window_delta_naive_d0.5_cpg5_meancov5/naive_windows_1000bp_delta_ge_0.5_cpg5_meancov5.tsv"))
n_neg <- read_dt(file.path(BASE, "window_delta_naive_d0.5_cpg5_meancov5/naive_windows_1000bp_delta_le_-0.5_cpg5_meancov5.tsv"))

t_pos <- read_dt(file.path(BASE, "window_delta_TSC_d0.5_cpg5_meancov5/TSC_windows_1000bp_delta_ge.tsv"))
t_neg <- read_dt(file.path(BASE, "window_delta_TSC_d0.5_cpg5_meancov5/TSC_windows_1000bp_delta_le.tsv"))

# =========================
# Barplot function
# =========================

make_bar <- function(p_pos, n_pos, t_pos, p_neg, n_neg, t_neg, title_suffix) {

  bar_dt <- rbind(
    data.table(group="h1 - h2 > 0.5", state="Prime", n=nrow(p_pos)),
    data.table(group="h1 - h2 > 0.5", state="Naive", n=nrow(n_pos)),
    data.table(group="h1 - h2 > 0.5", state="TSC",   n=nrow(t_pos)),

    data.table(group="h1 - h2 < -0.5", state="Prime", n=nrow(p_neg)),
    data.table(group="h1 - h2 < -0.5", state="Naive", n=nrow(n_neg)),
    data.table(group="h1 - h2 < -0.5", state="TSC",   n=nrow(t_neg))
  )

  bar_dt[, state := factor(state, levels=c("Prime","Naive","TSC"))]

  ggplot(bar_dt, aes(state, n, fill=state)) +
    geom_col(width=0.7) +
    geom_text(aes(label=n), vjust=-0.5, size=4) +
    scale_fill_manual(values=cols) +
    facet_wrap(~group, nrow=1) +
    labs(
      title=paste0("Allele-specific windows ", title_suffix),
      subtitle="1000 bp windows; CpG≥5, mean cov≥5",
      x=NULL,
      y="Number of windows"
    ) +
    theme_bw(base_size=12) +
    theme(
      legend.position="none",
      plot.title=element_text(face="bold")
    )
}

p_bar_withX <- make_bar(
  p_pos, n_pos, t_pos,
  p_neg, n_neg, t_neg,
  "(chr1-22 + chrX)"
)

p_bar_noX <- make_bar(
  keep_auto(p_pos), keep_auto(n_pos), keep_auto(t_pos),
  keep_auto(p_neg), keep_auto(n_neg), keep_auto(t_neg),
  "(chr1-22 only)"
)

p_bar_final <- p_bar_withX / p_bar_noX +
  plot_annotation(
    title="Allele-specific window counts with and without sex chromosomes",
    theme=theme(plot.title=element_text(size=16, face="bold"))
  )

ggsave(
  file.path(OUT, "barplot_counts_d0.5_with_and_without_X.pdf"),
  p_bar_final,
  width=9,
  height=8
)

# =========================
# Venn functions
# =========================

draw_three <- function(p_dt, n_dt, t_dt, panel_title, vp) {

  pushViewport(vp)

  grid.text(
    panel_title,
    y=unit(0.95, "npc"),
    gp=gpar(fontsize=13, fontface="bold")
  )

  pushViewport(viewport(y=0.45, height=0.85))

  p_set <- unique(p_dt$id)
  n_set <- unique(n_dt$id)
  t_set <- unique(t_dt$id)

  nP <- length(p_set)
  nN <- length(n_set)
  nT <- length(t_set)

  nPN  <- length(intersect(p_set, n_set))
  nPT  <- length(intersect(p_set, t_set))
  nNT  <- length(intersect(n_set, t_set))
  nPNT <- length(Reduce(intersect, list(p_set, n_set, t_set)))

  venn.plot <- draw.triple.venn(
    area1 = nP,
    area2 = nN,
    area3 = nT,
    n12 = nPN,
    n13 = nPT,
    n23 = nNT,
    n123 = nPNT,
    category = c("Prime","Naive","TSC"),
    fill = c(cols["Prime"], cols["Naive"], cols["TSC"]),
    cex = 1.0,
    cat.cex = 1.0,
    cat.pos = c(-20, 20, 180)
  )

  grid.draw(venn.plot)
  popViewport(2)
}

# =========================
# Venn PDF: with X and no X
# 2 rows x 2 columns
# =========================

pdf(file.path(OUT, "venn_3way_d0.5_with_and_without_X.pdf"), width=9, height=8)

grid.newpage()
grid.text(
  "3-way overlap of allele-specific windows",
  y=unit(0.985, "npc"),
  gp=gpar(fontsize=16, fontface="bold")
)

grid.text(
  "Top: chr1-22 + chrX; Bottom: chr1-22 only",
  y=unit(0.95, "npc"),
  gp=gpar(fontsize=11)
)

pushViewport(viewport(y=0.47, height=0.88, layout=grid.layout(2,2)))

draw_three(
  p_pos, n_pos, t_pos,
  "With X: h1 - h2 > 0.5",
  viewport(layout.pos.row=1, layout.pos.col=1)
)

draw_three(
  p_neg, n_neg, t_neg,
  "With X: h1 - h2 < -0.5",
  viewport(layout.pos.row=1, layout.pos.col=2)
)

draw_three(
  keep_auto(p_pos), keep_auto(n_pos), keep_auto(t_pos),
  "No X: h1 - h2 > 0.5",
  viewport(layout.pos.row=2, layout.pos.col=1)
)

draw_three(
  keep_auto(p_neg), keep_auto(n_neg), keep_auto(t_neg),
  "No X: h1 - h2 < -0.5",
  viewport(layout.pos.row=2, layout.pos.col=2)
)

dev.off()

message("[DONE] Outputs written to: ", OUT)
