# Function defining outbrekas
calculate_outbreaks <- function(result, all_base_compartments) {
  infected_compartments_string <- paste(all_base_compartments[grep("I", all_base_compartments)], collapse = "+")
  infected_compartments_expr <- parse(text = infected_compartments_string)
  
  infected <- data.frame(
    "infected" = result[, eval(infected_compartments_expr)],
    "node" = result[, "node"],
    "time" = result[, "time"]
  )
  
  max_infected <- infected %>%
    group_by(node) %>%
    summarise(max_infected = max(infected))
  
  transmission_after_121 <- infected %>%
    filter(time >= 121) %>%
    group_by(node) %>%
    summarise(transmission = any(infected > 0))
  
  outbreak_or_not <- transmission_after_121 %>%
    left_join(max_infected, by = "node") %>%
    mutate(is_outbreak = (transmission & max_infected >= 100)) # Criterion is 100 infected at the same time and infected after 121 months
  
  outbreak_subset <- outbreak_or_not %>% filter(is_outbreak)
  non_outbreak_subset <- outbreak_or_not %>% filter(!is_outbreak)
  
  outbreak_subset_results <- result[result$node %in% outbreak_subset$node, ]
  non_outbreak_subset_results <- result[result$node %in% non_outbreak_subset$node, ]
  
  # Relabel nodes
  unique_nodes <- sort(unique(outbreak_subset_results$node))
  translation_list <- setNames(seq_along(unique_nodes), unique_nodes)
  outbreak_subset_results$node <- translation_list[as.character(outbreak_subset_results$node)]
  
  unique_nodes <- sort(unique(non_outbreak_subset_results$node))
  translation_list <- setNames(seq_along(unique_nodes), unique_nodes)
  non_outbreak_subset_results$node <- translation_list[as.character(non_outbreak_subset_results$node)]
  
  return(list(
    outbreak_fraction = sum(outbreak_or_not$is_outbreak) / nrow(outbreak_or_not),
    outbreak_subset = outbreak_subset_results,
    non_outbreak_fraction = 1 - sum(outbreak_or_not$is_outbreak) / nrow(outbreak_or_not),
    non_outbreak_subset = non_outbreak_subset_results
  ))
}


#   ________________________________________________________________________________________
#   Function that reruns simulations to desired number of outbreaks and non outbreaks   ####


rerun_simulations_until_desired <- function(
    scenario_model, 
    all_base_compartments, 
    total_outbreaks_needed = 1000,
    total_non_outbreaks_needed = NULL,
    initial_results = NULL,
    set_threshold = 1000,
    print_beta = FALSE
) {
  outbreak_subset <- data.frame()
  non_outbreak_subset <- data.frame()
  
  total_outbreaks_gathered <- 0
  total_non_outbreaks_gathered <- 0
  total_simulations_processed <- 0
  total_outbreaks_observed <- 0
  total_non_outbreaks_observed <- 0
  
  while (total_outbreaks_gathered < total_outbreaks_needed ||
         (!is.null(total_non_outbreaks_needed) && total_non_outbreaks_gathered < total_non_outbreaks_needed)) {
    
    # Run the simulation
    result_tmp <- run(scenario_model)
    result_tmp <- setDT(trajectory(result_tmp))
    outbreak_results <- calculate_outbreaks(result_tmp, all_base_compartments)
    
    n_nodes_in_batch <- length(unique(result_tmp$node))
    total_simulations_processed <- total_simulations_processed + n_nodes_in_batch
    n_outbreaks_in_batch <- length(unique(outbreak_results$outbreak_subset$node))
    n_non_outbreaks_in_batch <- length(unique(outbreak_results$non_outbreak_subset$node))
    
    total_outbreaks_observed <- total_outbreaks_observed + n_outbreaks_in_batch
    total_non_outbreaks_observed <- total_non_outbreaks_observed + n_non_outbreaks_in_batch
    
    if (total_outbreaks_gathered < total_outbreaks_needed && n_outbreaks_in_batch > 0) {
      outbreaks_needed_now <- min(total_outbreaks_needed - total_outbreaks_gathered, n_outbreaks_in_batch)
      nodes_to_add <- unique(outbreak_results$outbreak_subset$node)[1:outbreaks_needed_now]
      subset_to_add <- outbreak_results$outbreak_subset[outbreak_results$outbreak_subset$node %in% nodes_to_add, ]
      
      # Adjust node labels for outbreaks
      if (nrow(outbreak_subset) > 0) {
        max_node <- max(outbreak_subset$node, na.rm = TRUE)
      } else {
        max_node <- 0
      }
      subset_to_add$node <- subset_to_add$node + max_node
      
      # Ensure that nodes from 1-100 are properly offset
      #print(paste0("Adjusted outbreak nodes (expected max_node offset): ", max_node))
      outbreak_subset <- rbind(outbreak_subset, subset_to_add)
      total_outbreaks_gathered <- total_outbreaks_gathered + length(unique(subset_to_add$node))
    }
    
    if (!is.null(total_non_outbreaks_needed) && total_non_outbreaks_gathered < total_non_outbreaks_needed && n_non_outbreaks_in_batch > 0) {
      non_outbreaks_needed_now <- min(total_non_outbreaks_needed - total_non_outbreaks_gathered, n_non_outbreaks_in_batch)
      nodes_to_add <- unique(outbreak_results$non_outbreak_subset$node)[1:non_outbreaks_needed_now]
      subset_to_add <- outbreak_results$non_outbreak_subset[outbreak_results$non_outbreak_subset$node %in% nodes_to_add, ]
      
      # Adjust node labels for non-outbreaks
      if (nrow(non_outbreak_subset) > 0) {
        max_node <- max(non_outbreak_subset$node, na.rm = TRUE)
      } else {
        max_node <- 0
      }
      subset_to_add$node <- subset_to_add$node + max_node
      
      # Ensure that nodes from 1-100 are properly offset
      #print(paste0("Adjusted non-outbreak nodes (expected max_node offset): ", max_node))
      non_outbreak_subset <- rbind(non_outbreak_subset, subset_to_add)
      total_non_outbreaks_gathered <- total_non_outbreaks_gathered + length(unique(subset_to_add$node))
    }
    
    print(paste0("Outbreaks gathered: ", total_outbreaks_gathered, " / ", total_outbreaks_needed))
    if (!is.null(total_non_outbreaks_needed)) {
      print(paste0("Non-outbreaks gathered: ", total_non_outbreaks_gathered, " / ", total_non_outbreaks_needed))
    }
  }
  
  total_simulations_observed <- total_outbreaks_observed + total_non_outbreaks_observed
  outbreak_fraction <- total_outbreaks_observed / total_simulations_observed
  non_outbreak_fraction <- total_non_outbreaks_observed / total_simulations_observed
  
  return(list(
    outbreak_subset = outbreak_subset,
    outbreak_fraction = outbreak_fraction,
    non_outbreak_subset = non_outbreak_subset,
    non_outbreak_fraction = non_outbreak_fraction
  ))
}