# 02_NMDS_plot.R
source("scripts/utils_packages.R")
df <- readr::read_csv("data/prey_data.csv", show_col_types = FALSE)
names(df) <- tolower(names(df))

if ("prey" %in% names(df)) {
  bad_levels <- tolower(c("unidentified","nieoznaczone","niezidentyfikowane","unknown","na"))
  dat_wide <- df %>%
    dplyr::mutate(prey = stringr::str_trim(tolower(as.character(prey)))) %>%
    dplyr::filter(!(prey %in% bad_levels)) %>%
    dplyr::count(id, species, prey, name = "n_items") %>%
    tidyr::pivot_wider(names_from = prey, values_from = n_items, values_fill = 0)
} else {
  taxa_possible <- c("diptera","odonata","trichoptera","orthoptera","coleoptera",
                     "ephemeroptera","pisces","lepidoptera","hymenoptera","unidentified")
  taxa_cols <- intersect(taxa_possible, names(df))
  dat_wide <- df %>%
    dplyr::group_by(id, species) %>%
    dplyr::summarise(dplyr::across(dplyr::all_of(taxa_cols),
                     ~ sum(suppressWarnings(as.numeric(.)), na.rm = TRUE)), .groups = "drop")
}
dat_wide <- dat_wide %>% dplyr::select(-dplyr::any_of("unidentified"))
meta_cols <- c("id","species")
prey_cols <- setdiff(names(dat_wide), meta_cols)

prey_mat_rel <- dat_wide %>%
  dplyr::mutate(total_prey = pmax(1, rowSums(dplyr::across(dplyr::all_of(prey_cols)), na.rm = TRUE))) %>%
  dplyr::mutate(dplyr::across(dplyr::all_of(prey_cols), ~ .x / total_prey)) %>%
  dplyr::select(dplyr::all_of(prey_cols)) %>%
  as.matrix()

set.seed(42)
nmds <- vegan::metaMDS(prey_mat_rel, distance = "bray", k = 2, trymax = 100, autotransform = FALSE)
scr <- as.data.frame(vegan::scores(nmds))
scr$id <- dat_wide$id
scr$species <- factor(dat_wide$species, levels = c("leucopterus","niger"))
readr::write_csv(scr, "outputs/NMDS_scores.csv")

get_hull <- function(d) d[chull(d$NMDS1, d$NMDS2), c("NMDS1","NMDS2")]
hulls <- scr %>% dplyr::group_by(species) %>% dplyr::reframe(get_hull(cur_data()), .id = NULL)

fit <- vegan::envfit(nmds, prey_mat_rel, permutations = 0)
tax <- as.data.frame(fit$vectors$arrows); tax$taxon <- rownames(tax)

p <- ggplot2::ggplot(scr, ggplot2::aes(x = NMDS1, y = NMDS2)) +
  ggplot2::geom_point(ggplot2::aes(shape = species), size = 3) +
  ggplot2::geom_polygon(data = hulls, ggplot2::aes(group = species), fill = NA, linewidth = 1) +
  ggplot2::geom_segment(data = tax, ggplot2::aes(x = 0, y = 0, xend = NMDS1, yend = NMDS2),
                        arrow = ggplot2::arrow(length = grid::unit(0.15, "cm"))) +
  ggrepel::geom_text_repel(data = tax, ggplot2::aes(x = NMDS1, y = NMDS2, label = taxon)) +
  ggplot2::labs(title = paste0("NMDS (stress = ", round(nmds$stress, 3), ")"),
                x = "NMDS1", y = "NMDS2") +
  ggplot2::theme_minimal(base_size = 12)

ggplot2::ggsave("outputs/figures/Fig2_NMDS_plot.png", p, width = 7, height = 5, dpi = 300)
