library(dplyr)
library(forcats)
library(ggplot2)
library(gridExtra)
library(paletteer)
library(rcartocolor)

source("./R/utils.R")

combine_imputed <- read.csv("./data/trait_data_imputed.csv")
combine_reported <- read.csv("./data/trait_data_reported.csv")
filters <- readRDS("./data/Sppdb_list_100replicates.rds")
filters_ind <- readRDS("./data/IND_Sppdb_list_100replicates.rds")

# Definition of size categories (g)
size_cat_breaks <- c(0, 100, 1000, 10000, 45000, 90000, 180000, 360000, Inf)

sel_traits <- combine_reported[, c(1:7, 45)]
incomplete <- is.na(sel_traits$adult_mass_g) |
  is.na(sel_traits$trophic_level)

sel_traits_complete <- sel_traits[!incomplete, ]
filter_sp_names <- filters[[1]][[1]][["Spp_names"]]
# When there are rows with the same IUCN name, select the one that has
# the same specific name (epithet) in the phylacine_binomial column
not_recognised <- sel_traits_complete$iucn2020_binomial %in% "Not recognised"
sel_traits_complete <- sel_traits_complete[!not_recognised, ]
dup_rows <- duplicated(sel_traits_complete$iucn2020_binomial)
dup_names <- unique(sel_traits_complete$iucn2020_binomial[dup_rows])
not_dup_rows <- !sel_traits_complete$iucn2020_binomial %in% dup_names
dphy_specific_name <- sapply(
  strsplit(sel_traits_complete$phylacine_binomial, split = " "),
  function(x) x[2]
)
not_dup_rows <- not_dup_rows | sel_traits_complete$species == dphy_specific_name
sel_traits_complete <- sel_traits_complete[not_dup_rows, ]

# Probably no trait data
unidentified <- setdiff(filter_sp_names, sel_traits_complete$iucn2020_binomial)
# Probably no range data or marine mammals
no_filter <- setdiff(sel_traits_complete$iucn2020_binomial, filter_sp_names)

# It is necessary to exclude from the trait database the rows with
# names that are not considered in the filters so that the number of
# species that are lost after applying the filter is properly
# calculated
names_in_filters <- sel_traits_complete$iucn2020_binomial %in% filter_sp_names
sel_traits_complete <- sel_traits_complete[names_in_filters, ]

# Size categories
sel_traits_complete$size_cat <- cut(
  sel_traits_complete$adult_mass_g,
  breaks = size_cat_breaks
)

# Order categories
sel_traits_complete$order_cat <- as.factor(sel_traits_complete$order) |>
  fct_lump_min(90)

# Alternative size categories
size_cat_breaks_2 <- c(0, 100, 1000, 45000, Inf)
sel_traits_complete$size_cat2 <- cut(
  sel_traits_complete$adult_mass_g,
  breaks = size_cat_breaks_2
)

filtered_trait_data <- lapply(filters, function(filter_tables) {
  # Three tables: 25%, 50%, 75%
  lapply(filter_tables, function(filter_table) {
    traits_whole_fam <- filter_trait_table(
      "Taxonomy whole-fam",
      trait_table = sel_traits_complete,
      filter_table
    )
    traits_tax_weighted <- filter_trait_table(
      "Taxonomy weighted",
      trait_table = sel_traits_complete,
      filter_table
    )
    list(whole_fam = traits_whole_fam, tax_weighted = traits_tax_weighted)
  })
})

# Objects with filtered trait tables
# TODO: the `filtered_trait_data` object should be redundant now, but
# double check
filtered_traits_seq <- filters |>
  lapply(function(replicate) {
    lapply(replicate, function(sampling_perc) {
      filter_list <- lapply(
        names(sampling_perc)[-(1:3)],
        function(filter_name) {
          filter_trait_table(filter_name, sel_traits_complete, sampling_perc)
        }
      )
      names(filter_list) <- names(sampling_perc)[-(1:3)]
      filter_list
    })
  })

filtered_traits_ind <- filters_ind |>
  lapply(function(replicate) {
    lapply(replicate, function(sampling_perc) {
      filter_list <- lapply(
        names(sampling_perc)[-(1:3)],
        function(filter_name) {
          filter_trait_table(filter_name, sel_traits_complete, sampling_perc)
        }
      )
      names(filter_list) <- names(sampling_perc)[-(1:3)]
      filter_list
    })
  })

# Sequential filter plots
seq_filter25 <- taxonomy_by_filter("25", filters, sel_traits_complete)
seq_filter50 <- taxonomy_by_filter("50", filters, sel_traits_complete)
seq_filter75 <- taxonomy_by_filter("75", filters, sel_traits_complete)

# Independent filters
ind_filter25 <- taxonomy_by_filter("25", filters_ind, sel_traits_complete)
ind_filter50 <- taxonomy_by_filter("50", filters_ind, sel_traits_complete)
ind_filter75 <- taxonomy_by_filter("75", filters_ind, sel_traits_complete)

filter_df_list <- list(
  filter25 = seq_filter25,
  filter50 = seq_filter50,
  filter75 = seq_filter75
)

indfilter_df_list <- list(
  indfilter25 = ind_filter25,
  indfilter50 = ind_filter50,
  indfilter75 = ind_filter75
)

filt_tax_path <- "./plots/filter_taxonomic_composition"
dir.create(filt_tax_path, recursive = TRUE)

# For some reason, the palette provided in the package paletteer for
# rcartocolor::Safe is different, with two of its colours being rather
# similar, so the original is used here
op <- options(
  ggplot2.discrete.fill = rcartocolor::carto_pal(10, "Safe"),
  ggplot2.discrete.colour = rcartocolor::carto_pal(10, "Safe")
)

for (filt_name in names(filter_df_list)) {
  tax_filter_df <- filter_df_list[[filt_name]]
  tax_filter_df <- filter(tax_filter_df, filter != "Taxonomy whole-fam")
  # Create line plots of retained counts and percentages per order
  perc_taxw <- ggplot(tax_filter_df) +
    geom_ribbon(
      aes(
        x = filter, ymin = perc_q05, ymax = perc_q95,
        fill = order, group = order
      ),
      alpha = 0.2
    ) +
    geom_line(aes(filter, perc_mean, col = order, group = order)) +
    xlab("Filter") +
    ylab("Retained species (%)") +
    labs(col = "Order", fill = "Order") +
    theme_classic() +
    theme(axis.text.x = element_text(angle = 45, hjust = 1))

  cnt_taxw <- ggplot(tax_filter_df) +
    geom_ribbon(
      aes(
        x = filter, ymin = count_q05, ymax = count_q95,
        fill = order, group = order
      ),
      alpha = 0.2
    ) +
    geom_line(aes(filter, count_mean, col = order, group = order)) +
    xlab("Filter") +
    ylab("Retained species") +
    labs(col = "Order", fill = "Order") +
    theme_classic() +
    theme(axis.text.x = element_text(angle = 45, hjust = 1))

  # Logarithm
  cnt_taxw_log <- ggplot(tax_filter_df) +
    geom_line(aes(filter, count_mean, col = order, group = order)) +
    scale_y_log10() +
    xlab("Filter") +
    ylab("Retained species (logarithmic scale)") +
    labs(col = "Order", fill = "Order") +
    theme_classic() +
    theme(axis.text.x = element_text(angle = 45, hjust = 1))

  cnt_barplot <- ggplot(tax_filter_df) +
    geom_bar(aes(filter, weight = count_mean, fill = order, group = order)) +
    xlab("Filter") +
    ylab("Species count") +
    labs(fill = "Order") +
    theme_classic()
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_count_barplot.svg")),
    cnt_barplot,
    width = 12, height = 7
  )

  perc_barplot <- ggplot(tax_filter_df) +
    geom_bar(
      aes(filter, weight = filt_total_perc, fill = order, group = order)
    ) +
    xlab("Filter") +
    ylab("Species count (%)") +
    labs(fill = "Order") +
    theme_classic()
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_perc_barplot.svg")),
    perc_barplot,
    width = 12, height = 7
  )

  svg(
    file.path(filt_tax_path, paste0(filt_name, ".svg")),
    width = 24, height = 6,
    pointsize = 16
  )
  grid.arrange(perc_taxw, cnt_taxw, cnt_taxw_log, nrow = 1)
  dev.off()
}

# SVG files to manipulate in a vector graphics editor to create the
# final figure
for (filt_name in names(indfilter_df_list)) {
  tax_filter_df <- indfilter_df_list[[filt_name]]
  cnt_barplot <- ggplot(tax_filter_df) +
    geom_bar(aes(filter, weight = count_mean, fill = order, group = order))
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_count_barplot.svg")),
    cnt_barplot,
    width = 12, height = 7
  )
  perc_barplot <- ggplot(tax_filter_df) +
    geom_bar(aes(filter, weight = filt_total_perc,
                 fill = order, group = order)) +
    scale_colour_discrete()
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_perc_barplot.svg")),
    perc_barplot,
    width = 12, height = 7
  )
}

# Total counts (n and mean n), needed for the figures that are done manually
for (filt_name in names(indfilter_df_list)) {
  filt_df <- indfilter_df_list[[filt_name]]
  message(filt_name)
  message("====================")
  print(tapply(filt_df$count_mean, filt_df$filter, sum))
  message("\n")
}

# Retained percentage of species by order and independent filter
selected_filters <- names(filters[[1]][[1]])[-c(1, 2, 3, 9)]

order_count_ind <- lapply(c("25", "50", "75"), function(sampling_perc) {
  filtered_traits_ind |>
    lapply(function(replicate) {
      tbl_list <- replicate[[sampling_perc]]
      tbl_list <- tbl_list[selected_filters]
      lapply(tbl_list, function(tbl) table(tbl$order_cat))
    })
})
names(order_count_ind) <- c("25", "50", "75")

# Calculate retained percentage by order
order_retained_ind <- order_count_ind |>
  lapply(function(sampling_perc) {
    lapply(sampling_perc, function(replicate) {
      lapply(replicate, function(f) f / replicate$Land)
    })
  })

# Convert to a data frame for convenience
retained_perc <- order_retained_ind |>
  lapply(function(sampling_perc) {
    df_list <- lapply(sampling_perc, function(replicate) {
      rep_lst <- lapply(replicate, as.data.frame, stringsAsFactors = FALSE)
      rep_df <- do.call(
        rbind, unname(Map(cbind, filter = names(rep_lst), rep_lst))
      )
      names(rep_df) <- c("filter", "order_cat", "ret_perc")
      rep_df
    })
    do.call(rbind, unname(Map(cbind, seed = names(df_list), df_list)))
  }) |>
  lapply(function(sampling_perc) {
    # Convert order column to a factor with the same levels as the trait
    # data frame to keep plots consistent
    sampling_perc$order_cat <- factor(
      sampling_perc$order_cat,
      levels = levels(sel_traits_complete$order_cat)
    )
    # Remove "Land" rows, since these are the original data before
    # applying filters
    sampling_perc <- sampling_perc[sampling_perc$filter != "Land", ]
    sampling_perc$filter <- factor(
      sampling_perc$filter,
      levels = c("UcS", "Range size", "Body size",
                 "Human spatial", "Taxonomy weighted")
    )
    sampling_perc
  })

options(
  ggplot2.discrete.fill = fill_alpha(rcartocolor::carto_pal(10, "Safe"), 0.7)
)

# Effect of each filter across orders (percentage of retained species)
for (perc in names(retained_perc)) {
  d <- retained_perc[[perc]]
  filt_name <- paste0("indfilter", perc)
  # Control order of boxplots
  perc_boxplot <- ggplot(d) +
    geom_boxplot(
      aes(ret_perc, order_cat, fill = order_cat, col = order_cat),
      position = position_dodge(1),
      outlier.size = 0.5
    ) +
    facet_grid(rows = vars(filter)) +
    scale_y_discrete(limits = rev(levels(d$order_cat))) +
    xlim(0, 1) +
    xlab("Retained species (%)") +
    labs(fill = "Order", col = "Order") +
    theme_classic() +
    theme(
      panel.border = element_rect(colour = "black", fill = NA),
      axis.text.y = element_blank(),
      axis.ticks.y = element_blank(),
      axis.title.y = element_blank()
    )
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_retained_boxplot_filt.svg")),
    perc_boxplot,
    width = 10, height = 10
  )
}

# Same information as the plots above, but boxplots are grouped by
# order instead of by filter
for (perc in names(retained_perc)) {
  d <- retained_perc[[perc]]
  filt_name <- paste0("indfilter", perc)
  # Control order of boxplots
  perc_boxplot <- ggplot(d) +
    geom_boxplot(
      aes(filter, ret_perc, fill = order_cat, col = order_cat),
      position = position_dodge(1),
      outlier.size = 0.5
    ) +
    facet_wrap(vars(order_cat), nrow = 2) +
    ylim(0, 1) +
    xlab("Filter") +
    ylab("Retained species (%)") +
    labs(fill = "Order", col = "Order") +
    theme_classic() +
    theme(
      panel.border = element_rect(colour = "black", fill = NA),
      axis.text.x = element_text(angle = 45, hjust = 1)
    )
  ggsave(
    file.path(filt_tax_path, paste0(filt_name, "_retained_boxplot.svg")),
    perc_boxplot,
    width = 16, height = 9
  )
}

# More complicated heatmaps
cat_count <- category_count_table(
  sel_traits_complete, "size_cat2", "trophic_level"
)
dimnames(cat_count)[[2]] <- c("herbivore", "omnivore", "carnivore")

# Taxonomic composition per category
categ_taxon <- tapply(
  sel_traits_complete$order_cat,
  list(sel_traits_complete$size_cat2, sel_traits_complete$trophic_level),
  table
)

cat_combs <- expand.grid(rownames(cat_count), 1:3, stringsAsFactors = FALSE)
names(cat_combs) <- c("size_cat2", "trophic_level")

filter_combs <- expand.grid(
  perc = c("25", "50", "75"),
  tax = c("whole_fam", "tax_weighted"),
  stringsAsFactors = FALSE
)

order_comp_all <- vector(mode = "list", length = nrow(filter_combs))
for (i in seq_len(nrow(filter_combs))) {
  comb_row <- filter_combs[i, ]
  perc <- comb_row$perc
  tax <- comb_row$tax
  comp_list <- lapply(seq_len(nrow(cat_combs)), function(row_idx) {
    cat_comb_row <- cat_combs[row_idx, ]
    size <- cat_comb_row$size_cat
    diet <- cat_comb_row$trophic_level
    counts_100 <- sapply(filtered_trait_data, function(trait_df_list) {
      trait_df <- trait_df_list[[perc]][[tax]]
      sel <- trait_df$size_cat2 == size & trait_df$trophic_level == diet
      trait_df_sel <- trait_df[sel, ]
      counts <- as.numeric(table(trait_df_sel$order_cat))
      names(counts) <- levels(trait_df_sel$order_cat)
      counts
    }) |> t()
    counts_100
  })
  names(comp_list) <- paste(
    cat_combs$size_cat2,
    rep(c("herbivore", "omnivore", "carnivore"), each = 4)
  )
  order_comp_all[[i]] <- comp_list
}
names(order_comp_all) <- paste(
  rep(c("Taxonomic whole-families", "Taxonomic weighted"), each = 3),
  rep(c("25%", "50%", "75%"), 2)
)

order_comp_count_means <- rapply(order_comp_all, colMeans, how = "replace")

filter_cat_count_diet <- category_count_tables(
  filtered_trait_data, "size_cat2", "trophic_level"
)

species_ret_perc_lst <- rapply(filter_cat_count_diet, function(count_tbl) {
  perc_tbl <- count_tbl / cat_count
  perc_tbl
}, how = "replace")
species_ret_perc_arr <- collect_tables(species_ret_perc_lst)
ret_perc_mean_list <- lapply(seq_len(nrow(filter_combs)), function(row_idx) {
  combn_row <- filter_combs[row_idx, ]
  perc <- combn_row$perc
  tax_filter <- combn_row$tax
  apply(species_ret_perc_arr[[perc]][[tax_filter]], c(1, 2), mean)
})
names(ret_perc_mean_list) <- paste(filter_combs$perc, filter_combs$tax)

# Plots with too much information
n_cat_x <- 3
n_cat_y <- 4
x_coords <- seq(0, 1, length.out = n_cat_x)
y_coords <- seq(0, 1, length.out = n_cat_y)
# There should be more than one category for this to work
cat_x_width <- x_coords[2] - x_coords[1]
cat_y_width <- y_coords[2] - y_coords[1]

rect_width <- cat_x_width * (2 / 3)
rect_height <- cat_y_width * (2 / 3)
orders <- levels(sel_traits_complete$order_cat)
n_orders <- length(orders)

col_pal <- carto_pal(n_orders, "Safe")
ret_pal <- paletteer_c("viridis::mako", 20)

plot_dir <- "./plots/heatmaps"
for (filter_idx in seq_len(nrow(filter_combs))) {
  after_filter_mean <- order_comp_count_means[[filter_idx]]
  main_title <- names(order_comp_count_means)[[filter_idx]]

  png(
    file.path(
      plot_dir,
      paste0(paste0(filter_combs[filter_idx, ], collapse = "_"), ".png")
    ),
    width = 1920,
    height = 1920,
    res = 72,
    pointsize = 32
  )
  layout(cbind(c(1, 1), c(2, 3)), widths = c(0.9, 0.15))
  image(
    t(ret_perc_mean_list[[filter_idx]]),
    col = ret_pal,
    axes = FALSE,
    zlim = c(0, 1),
    main = main_title
  )
  axis(
    1,
    at = seq(0, 1, length.out = n_cat_x),
    labels = c("Herbivore", "Omnivore", "Carnivore"),
    tick = FALSE
  )
  axis(
    2,
    at = seq(0, 1, length.out = n_cat_y),
    labels = rownames(categ_taxon),
    tick = FALSE
  )
  mtext("Diet", 1, line = 2.5)
  mtext("Body mass", 2, line = 2.5)
  for (col in seq_len(ncol(categ_taxon))) {
    for (row in seq_len(nrow(categ_taxon))) {
      x <- x_coords[[col]]
      y <- y_coords[[row]]
      counts <- categ_taxon[row, col][[1]]
      percs <- counts / sum(counts)
      rect_y_coords <- c(0, cumsum(percs))
      after_filter_counts <- after_filter_mean[[(col - 1) * n_cat_y + row]]
      after_filter_percs <- after_filter_counts / sum(after_filter_counts)
      rect2_y_coords <- c(0, cumsum(after_filter_percs))
      x_offset <- rect_width / 2
      y_offset <- rect_height / 2
      rect_x_min <- x - x_offset
      rect_x_max <- x + x_offset
      rect_y_min <- y - y_offset
      rect_y_max <- y + y_offset
      bar_width <- rect_width * 0.45
      # All percentage vectors should have length > 1
      for (i in seq_len(length(rect_y_coords) - 1)) {
        rect(
          rect_x_min,
          rect_y_min + rect_height * rect_y_coords[[i]],
          rect_x_min + bar_width,
          rect_y_min + rect_height * rect_y_coords[[i + 1]],
          col = col_pal[i], border = NA
        )
      }
      bar2_x_min <- rect_x_min + rect_width * 0.55
      for (i in seq_len(length(rect2_y_coords) - 1)) {
        rect(
          bar2_x_min,
          rect_y_min + rect_height * rect2_y_coords[[i]],
          bar2_x_min + bar_width,
          rect_y_min + rect_height * rect2_y_coords[[i + 1]],
          col = col_pal[i], border = NA
        )
      }
      rect(
        rect_x_min, rect_y_min, rect_x_min + bar_width, rect_y_max,
        col = NA, border = "white"
      )
      rect(
        bar2_x_min, rect_y_min, rect_x_max, rect_y_max,
        col = NA, border = "white"
      )
      legend(
        rect_x_min + rect_width * 0.45 / 2, rect_y_min,
        paste0("n = ", sum(counts)),
        box.col = "white", bg = "white",
        x.intersp = 0, y.intersp = 0.2,
        adj = 0.1, xjust = 0.5, yjust = 1.25
      )
      mean_n_after <- round(sum(after_filter_counts), 1)
      legend(
        rect_x_min + rect_width * (0.55 + (0.45 / 2)), rect_y_min,
        bquote(bar(n) == .(mean_n_after)),
        box.col = "white", bg = "white",
        x.intersp = 0, y.intersp = 0.2,
        adj = 0.1, xjust = 0.5, yjust = 1.25
      )
    }
  }
  opar <- par(mar = c(0, 0, 0, 0))
  plot.new()
  legend(
    "center",
    legend = orders,
    col = NA,
    pt.bg = col_pal,
    pch = 22, pt.cex = 3,
    y.intersp = 2,
    bty = "n"
  )
  lg_rast <- as.raster(matrix(rev(ret_pal), ncol = 1))
  par(mar = c(3, 0, 2, 5))
  plot(c(0, 1), c(0, 1), type = "n", axes = FALSE, xlab = NA, ylab = NA)
  rasterImage(lg_rast, 0, 0, 1, 1)
  mtext("Retention %", 3, adj = 0.3)
  axis(4, at = seq(0, 1, 0.1), labels = seq(0, 100, 10), las = 1)
  par(opar)
  dev.off()
}

# Shannon entropy
shannon <- function(x) {
  p <- x / sum(x)
  -sum(p * log(p), na.rm = TRUE)
}

# 50% filter strength
orig <- categ_taxon["(0,100]", 1][[1]]
after_filter <- order_comp_count_means[[5]][[1]]
apply(categ_taxon, c(1, 2), function(x) shannon(x[[1]]))
sapply(order_comp_count_means[[5]], shannon)

options(op)
