library(ape)
library(ggtree)
library(ggforce)
library(phangorn)
library(ggnewscale)
library(viridisLite)

cladelab.df.getter <- function(gradlower, gradupper, labxlower, labxupper, labylower, labyupper) {
  x_steps <- seq(gradlower, gradupper, length.out = 201)
  alpha_steps <- seq(0, 1, length.out = 200)
  rect_data <- data.frame(xmin = x_steps[-201], xmax = x_steps[-1], 
                          ymin = labylower + 0.04, ymax = labyupper - 0.04,
                          alpha = alpha_steps)
  
  # bottom left, bottom right, top right, top left
  shape_data <- data.frame(x = c(labxlower, labxupper, labxupper, labxlower),
                           y = c(labylower, labylower, labyupper, labyupper))
  
  return(list(rect = rect_data, shape = shape_data))
}

visualize.tree <- function(inpath, saur_spec, orni_spec, ther_spec, index = NULL, brlen = T) {
  dir <- "~/Grive/Slater_Lab/Ornithoscelida_Ashley/"
  trees <- read.tree(paste0(dir, inpath))
  if (class(trees) == "multiPhylo") {
    tr <- trees[[index]]
  } else {
    tr <- trees
  }

  # This is for Baron et al.
  if ("Lewisuchus_et_Pseudolagosuchus" %in% tr$tip.label) {
    # Rename tip labels
    tr$tip.label[which(tr$tip.label == "Lewisuchus_et_Pseudolagosuchus")] <- "Lewisuchus/Pseudolagosuchus"
    tr$tip.label[which(tr$tip.label == "Massospondylus_kaalea")] <- "Massospondylus_kaalae"
    tr$tip.label[which(tr$tip.label == "Syntarsus_kayentakatae")] <- "\"Syntarsus\"_kayentakatae"
  
    # Clade anchors
    saur_spec_internal <- "Plateosaurus engelhardti"
    orni_spec_internal <- "Jeholosaurus shangyuanensis"
    ther_spec_internal <- "Dilophosaurus wetherilli"
  }
  
  # This is for Langer et al.
  if ("Lewisuchus" %in% tr$tip.label || "Lewisuchus_slash_Pseudolagosuchus" %in% tr$tip.label) {
    # Rename tip labels
    tr$tip.label[which(tr$tip.label %in% c("Lewisuchus", "Lewisuchus_slash_Pseudolagosuchus"))] <- "Lewisuchus/Pseudolagosuchus"
    tr$tip.label[which(tr$tip.label == "Syntarsus")] <- "\"Syntarsus\"_kayentakatae"
    
    # Clade anchors
    saur_spec_internal <- "Plateosaurus"
    orni_spec_internal <- "Jeholosaurus"
    ther_spec_internal <- "Dilophosaurus"
  }

  # Remove underscores from tip labels
  tr$tip.label <- sapply(strsplit(tr$tip.label, "_"), \(x) paste(x, collapse = " "))

  tr <- ladderize(tr)

  saur_mrca <- getMRCA(tr, c(saur_spec_internal, saur_spec))
  orni_mrca <- getMRCA(tr, c(orni_spec_internal, orni_spec))
  ther_mrca <- getMRCA(tr, c(ther_spec_internal, ther_spec))

  saur <- tr$tip.label[Descendants(tr, saur_mrca, type = "tips")[[1]]]
  orni <- tr$tip.label[Descendants(tr, orni_mrca, type = "tips")[[1]]]
  ther <- tr$tip.label[Descendants(tr, ther_mrca, type = "tips")[[1]]]

  saur_divg <- Ancestors(tr, saur_mrca, type = "parent")
  orni_divg <- Ancestors(tr, orni_mrca, type = "parent")
  ther_divg <- Ancestors(tr, ther_mrca, type = "parent")

  tree_plot_data <- ggtree(tr)$data

  saur_ylower <- min(tree_plot_data$y[tree_plot_data$label %in% saur]) - 0.6
  saur_yupper <- max(tree_plot_data$y[tree_plot_data$label %in% saur]) + 0.4
  saur_xlower <- tree_plot_data$x[which(tree_plot_data$node == saur_divg)]

  orni_ylower <- min(tree_plot_data$y[tree_plot_data$label %in% orni]) - 0.6
  orni_yupper <- max(tree_plot_data$y[tree_plot_data$label %in% orni]) + 0.4
  orni_xlower <- tree_plot_data$x[which(tree_plot_data$node == orni_divg)]

  ther_ylower <- min(tree_plot_data$y[tree_plot_data$label %in% ther]) - 0.6
  ther_yupper <- max(tree_plot_data$y[tree_plot_data$label %in% ther]) + 0.4
  ther_xlower <- tree_plot_data$x[which(tree_plot_data$node == ther_divg)]
  
  if (brlen) {
    xupper <- 1.23*max(tree_plot_data$x)
    xmax_global <- 1.25*max(tree_plot_data$x)
    saumorpha <- cladelab.df.getter(saur_xlower, xupper, xupper - 0.08, xupper + 0.1,
                                    saur_ylower, saur_yupper)
    ornischia <- cladelab.df.getter(orni_xlower, xupper, xupper - 0.08, xupper + 0.1,
                                    orni_ylower, orni_yupper)
    theropoda <- cladelab.df.getter(ther_xlower, xupper, xupper - 0.08, xupper + 0.1,
                                    ther_ylower, ther_yupper)
  } else {
    xupper <- 1.43*max(tree_plot_data$x)
    xmax_global <- 1.45*max(tree_plot_data$x)
    saumorpha <- cladelab.df.getter(saur_xlower, xupper, xupper - 0.8, xupper + 1,
                                    saur_ylower, saur_yupper)
    ornischia <- cladelab.df.getter(orni_xlower, xupper, xupper - 0.8, xupper + 1,
                                    orni_ylower, orni_yupper)
    theropoda <- cladelab.df.getter(ther_xlower, xupper, xupper - 0.8, xupper + 1,
                                    ther_ylower, ther_yupper)
  }

  clade_cols <- c(viridis(5)[2], viridis(3)[c(2, 3)])

  gg <- ggplot(tr) + xlab(NULL) + ylab(NULL) +
          coord_cartesian(xlim = c(0, xmax_global), clip = "off") +
          geom_shape(data = saumorpha$shape, aes(x = x, y = y), fill = clade_cols[1],
                     radius = unit(0.9, "cm")) +
          geom_rect(data = saumorpha$rect, aes(xmin = xmin, xmax = xmax, ymin = ymin, ymax = ymax,
                                               color = alpha, fill = alpha), show.legend = F) +
          scale_fill_gradient(low = "white", high = clade_cols[1]) +
          scale_colour_gradient(low = "white", high = clade_cols[1]) +
          new_scale_fill() +
          new_scale_color() +
          geom_shape(data = ornischia$shape, aes(x = x, y = y), fill = clade_cols[2],
                     radius = unit(0.9, "cm")) +
          geom_rect(data = ornischia$rect, aes(xmin = xmin, xmax = xmax, ymin = ymin, ymax = ymax,
                                               color = alpha, fill = alpha), show.legend = F) +
          scale_fill_gradient(low = "white", high = clade_cols[2]) +
          scale_colour_gradient(low = "white", high = clade_cols[2]) +
          new_scale_fill() +
          new_scale_color() +
          geom_shape(data = theropoda$shape, aes(x = x, y = y), fill = clade_cols[3],
                     radius = unit(0.9, "cm")) +
          geom_rect(data = theropoda$rect, aes(xmin = xmin, xmax = xmax, ymin = ymin, ymax = ymax,
                                               color = alpha, fill = alpha), show.legend = F) +
          scale_fill_gradient(low = "white", high = clade_cols[3]) +
          scale_colour_gradient(low = "white", high = clade_cols[3]) +
          geom_tree() + geom_tiplab(size = 4.5, fontface = "italic") +
          annotate(geom = "text", label = "Sauropodomorpha", x = xupper + 0.025,
                   y = (saur_ylower + saur_yupper)/2, angle = 270, size = 10, color = "white") +
          annotate(geom = "text", label = "Ornithischia", x = xupper + 0.025,
                   y = (orni_ylower + orni_yupper)/2, angle = 270, size = 10, color = "white") +
          annotate(geom = "text", label = "Theropoda", x = xupper + 0.025,
                   y = (ther_ylower + ther_yupper)/2, angle = 270, size = 10, color = "black") +
          theme(panel.border = element_blank(),
                panel.background = element_rect(fill = "white", colour = "white"),
                axis.text = element_blank(),
                axis.ticks = element_blank(),
                panel.grid.major.y = element_blank(),
                panel.grid.minor.y = element_blank(),
                panel.grid.major.x = element_blank(),
                panel.grid.minor.x = element_blank(),
                plot.margin = unit(c(0.05, 0, 0.05, 0.05), "cm"))
  
    if (!is.null(tr$node.label)) {
      to_exclude <- Ntip(tr) + 1
      # Very inelegant, but for some reason, we cannot stick to_exclude directly inside subset()
      if (to_exclude == 75) {
        gg <- gg + geom_nodelab(aes(subset = (node != 75), label = label), geom = "label",
                                size = 3, hjust = 0.75, label.padding = unit(0.15, "lines"))
      } else if (to_exclude == 84) {
        gg <- gg + geom_nodelab(aes(subset = (node != 84), label = label), geom = "label",
                                size = 3, hjust = 0.75, label.padding = unit(0.15, "lines"))
      }
    }
  
    if (brlen) {
      gg <- gg + geom_treescale(x = max(tree_plot_data$x), y = 5, width = 0.2, fontsize = 7) +
        annotate(geom = "text", label = "Expected substitutions per character",
                 x = max(tree_plot_data$x) + 0.1, y = 3, size = 4.5)
    }
  
    return(gg)
}

dir <- "~/Grive/Slater_Lab/Ornithoscelida_Ashley/"

pdf(paste0(dir, "Supp_figs/S1.pdf"), width = 15, height = 19)
visualize.tree("ML_unconstrained/baron_unconstrained.treefile",
               "Saturnalia tupiniquim", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S2.pdf"), width = 15, height = 19)
visualize.tree("ML_unconstrained/langer_unconstrained.treefile",
               "Thecodontosaurus", "Pisanosaurus", "Tawa")
dev.off()

pdf(paste0(dir, "Supp_figs/S3.pdf"), width = 15, height = 19)
visualize.tree("Tree_comparisons/BEA/baron_saurischia.treefile",
               "Staurikosaurus pricei", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S4.pdf"), width = 15, height = 19)
visualize.tree("Tree_comparisons/BEA/baron_ornithischiformes.treefile",
               "Saturnalia tupiniquim", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S5.pdf"), width = 15, height = 19)
visualize.tree("Tree_comparisons/LEA/langer_saurischia.treefile",
               "Eoraptor", "Pisanosaurus", "Tawa")
dev.off()

pdf(paste0(dir, "Supp_figs/S6.pdf"), width = 15, height = 19)
visualize.tree("Tree_comparisons/LEA/langer_ornithoscelida.treefile",
               "Agnosphitys", "Zupaysaurus", "Sarcosaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S7.pdf"), width = 15, height = 19)
visualize.tree("Exclude_decisive/Baron_sans_1_most_decisive_chars_partitions.nex.treefile",
               "Saturnalia tupiniquim", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S8.pdf"), width = 15, height = 19)
visualize.tree("Exclude_decisive/Baron_sans_5_most_decisive_chars_partitions.nex.treefile",
               "Saturnalia tupiniquim", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S9.pdf"), width = 15, height = 19)
visualize.tree("Exclude_decisive/Baron_sans_10_most_decisive_chars_partitions.nex.treefile",
               "Staurikosaurus pricei", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S10.pdf"), width = 15, height = 19)
visualize.tree("exclude_decisive/Baron_sans_all_most_decisive_chars_partitions.nex.treefile",
               "Staurikosaurus pricei", "Pisanosaurus mertii", "Eoraptor lunensis")
dev.off()

pdf(paste0(dir, "Supp_figs/S11.pdf"), width = 15, height = 19)
visualize.tree("exclude_decisive/Langer_sans_1_most_decisive_chars_partitions.nex.treefile",
               "Thecodontosaurus", "Pisanosaurus", "Chindesaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S12.pdf"), width = 15, height = 19)
visualize.tree("exclude_decisive/Langer_sans_5_most_decisive_chars_partitions.nex.treefile",
               "Thecodontosaurus", "Pisanosaurus", "Daemonosaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S13.pdf"), width = 15, height = 19)
visualize.tree("exclude_decisive/Langer_sans_10_most_decisive_chars_partitions.nex.treefile",
               "Thecodontosaurus", "Pisanosaurus", "Guaibasaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S14.pdf"), width = 15, height = 19)
visualize.tree("exclude_decisive/Langer_sans_all_most_decisive_chars_partitions.nex.treefile",
               "Thecodontosaurus", "Pisanosaurus", "Daemonosaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S15.pdf"), width = 15, height = 19)
visualize.tree("one_at_a_time/LEA_with_char_77_recoded_partitions.nex.treefile",
               "Eoraptor", "Staurikosaurus", "Daemonosaurus")
dev.off()

pdf(paste0(dir, "Supp_figs/S16.pdf"), width = 15, height = 19)
visualize.tree("one_at_a_time/LEA_with_char_148_recoded_partitions.nex.treefile",
               "Eoraptor", "Liliensternus", "Dracoraptor")
dev.off()

pdf(paste0(dir, "Supp_figs/S17.pdf"), width = 15, height = 19)
visualize.tree("one_at_a_time/LEA_with_char_363_recoded_partitions.nex.treefile",
               "Eoraptor", "Sarcosaurus", "Liliensternus")
dev.off()

pdf(paste0(dir, "Supp_figs/S18.pdf"), width = 15, height = 19)
visualize.tree("one_at_a_time/LEA_with_char_370_recoded_partitions.nex.treefile",
               "Eoraptor", "Lophostropheus", "Coelophysis")
dev.off()