# ============================================================
# SEMANTIC MAPS: MINIMAL GRAPH INFERENCE
# Requires two Excel files with sheets: semantic_counts, connections
# ============================================================

# ---- packages ---------------------------------------------------------------
pkgs <- c("openxlsx","readxl","dplyr","tidyr","stringr","purrr","tibble","igraph")
miss <- pkgs[!pkgs %in% rownames(installed.packages())]
if (length(miss)) install.packages(miss)
invisible(lapply(pkgs, library, character.only = TRUE))

# ---- file names (assumed to be in working directory) ------------------------
map1_file <- "semantic_maps_1.xlsx"
map2_file <- "semantic_maps_2.xlsx"
output_file <- "semantic_map_network_analysis.xlsx"

# ---- semantic functions considered "on the map" ----------------------------
on_map_functions <- c(
  "specific known", "general existential", "specific unknown",
  "existential quantification", "free choice", "random choice",
  "universal quantification"
)

# ---- analysis settings ------------------------------------------------------
use_binary_profiles <- TRUE   # presence/absence rather than raw frequencies
max_added_edges     <- 100L   # upper bound for edge augmentation

# ---- helper functions -------------------------------------------------------
clean_text <- function(x) {
  stringr::str_squish(stringr::str_replace_all(as.character(x), "\u00A0", " "))
}

find_column <- function(data, candidates, required = TRUE) {
  avail <- names(data)
  norm_avail <- tolower(clean_text(avail))
  norm_cand  <- tolower(clean_text(candidates))
  hit <- match(norm_cand, norm_avail, nomatch = 0)
  hit <- hit[hit > 0]
  if (length(hit) == 0) {
    if (!required) return(NA_character_)
    stop("No matching column found. Expected: ", paste(candidates, collapse = ", "),
         "\nFound: ", paste(avail, collapse = ", "))
  }
  avail[hit[[1]]]
}

read_sheet <- function(file, sheet) {
  if (!file.exists(file)) stop("File not found: ", file)
  if (!sheet %in% readxl::excel_sheets(file)) stop("Sheet '", sheet, "' missing in ", file)
  readxl::read_excel(file, sheet = sheet) |> tibble::as_tibble()
}

numeric_value <- function(x) {
  suppressWarnings(as.numeric(stringr::str_replace_all(as.character(x), ",", ".")))
}

# ---- data import ------------------------------------------------------------
read_dataset <- function(file, dataset_name) {
  message("Reading: ", dataset_name)
  sem_raw <- read_sheet(file, "semantic_counts")
  con_raw <- read_sheet(file, "connections")
  
  sem_lang  <- find_column(sem_raw, c("language", "lang"))
  sem_ser   <- find_column(sem_raw, c("series_id", "series"))
  sem_disp  <- find_column(sem_raw, c("display_name", "series_name", "modifier"), FALSE)
  sem_func  <- find_column(sem_raw, c("semantic_function", "semantic_standard", "semantic", "semantic_clean"))
  sem_cnt   <- find_column(sem_raw, c("count", "n", "frequency", "freq", "size", "value"))
  
  sem_counts <- sem_raw |>
    transmute(
      dataset = dataset_name,
      language = clean_text(.data[[sem_lang]]),
      series_id = clean_text(.data[[sem_ser]]),
      display_name = if (is.na(sem_disp)) clean_text(.data[[sem_ser]]) else clean_text(.data[[sem_disp]]),
      semantic_function = clean_text(.data[[sem_func]]),
      count = numeric_value(.data[[sem_cnt]])
    ) |>
    filter(!is.na(language), language != "",
           !is.na(series_id), series_id != "",
           !is.na(semantic_function), semantic_function != "",
           !is.na(count), count > 0) |>
    mutate(series_key = paste(language, series_id, sep = "::")) |>
    group_by(dataset, language, series_id, display_name, series_key, semantic_function) |>
    summarise(count = sum(count), .groups = "drop")
  
  con_lang  <- find_column(con_raw, c("language", "lang"))
  con_ser   <- find_column(con_raw, c("series_id", "series"))
  con_disp  <- find_column(con_raw, c("display_name", "series_name", "modifier"), FALSE)
  con_func  <- find_column(con_raw, c("semantic_function", "semantic_standard", "semantic", "semantic_clean"))
  con_ctx   <- find_column(con_raw, c("context", "context_standard", "context_clean"))
  con_cnt   <- find_column(con_raw, c("count", "n", "frequency", "freq", "size", "value"), FALSE)
  
  if (is.na(con_cnt)) {
    con_raw$count_auto <- 1
    con_cnt <- "count_auto"
  }
  
  connections <- con_raw |>
    transmute(
      dataset = dataset_name,
      language = clean_text(.data[[con_lang]]),
      series_id = clean_text(.data[[con_ser]]),
      display_name = if (is.na(con_disp)) clean_text(.data[[con_ser]]) else clean_text(.data[[con_disp]]),
      semantic_function = clean_text(.data[[con_func]]),
      context = clean_text(.data[[con_ctx]]),
      count = numeric_value(.data[[con_cnt]])
    ) |>
    filter(!is.na(language), language != "",
           !is.na(series_id), series_id != "",
           !is.na(semantic_function), semantic_function != "",
           !is.na(context), context != "",
           !is.na(count), count > 0) |>
    mutate(series_key = paste(language, series_id, sep = "::"),
           series_context_key = paste(language, series_id, context, sep = "::")) |>
    group_by(dataset, language, series_id, display_name, series_key, series_context_key, semantic_function, context) |>
    summarise(count = sum(count), .groups = "drop")
  
  list(semantic_counts = sem_counts, connections = connections)
}

# ---- profile matrix ---------------------------------------------------------
build_profile_matrix <- function(dataset, use_contexts, include_offmap) {
  if (use_contexts) {
    long <- dataset$connections |>
      select(semantic_function, feature = series_context_key, count)
  } else {
    long <- dataset$semantic_counts |>
      select(semantic_function, feature = series_key, count)
  }
  if (!include_offmap) long <- filter(long, semantic_function %in% on_map_functions)
  long <- long |>
    group_by(semantic_function, feature) |>
    summarise(value = sum(count), .groups = "drop")
  if (use_binary_profiles) long <- mutate(long, value = as.integer(value > 0))
  wide <- long |>
    pivot_wider(names_from = feature, values_from = value, values_fill = 0) |>
    arrange(semantic_function)
  if (nrow(wide) < 2) stop("Too few functions for a network.")
  mat <- as.matrix(wide |> select(-semantic_function))
  rownames(mat) <- wide$semantic_function
  storage.mode(mat) <- "numeric"
  mat
}

# ---- Jaccard similarity matrix ----------------------------------------------
jaccard_matrix <- function(x) {
  present <- x > 0
  fns <- rownames(x)
  n <- nrow(x)
  sim <- matrix(0, n, n, dimnames = list(fns, fns))
  for (i in seq_len(n)) {
    for (j in i:n) {
      union_n <- sum(present[i, ] | present[j, ])
      inter_n <- sum(present[i, ] & present[j, ])
      value <- if (union_n == 0) 0 else inter_n / union_n
      sim[i, j] <- value
      sim[j, i] <- value
    }
  }
  diag(sim) <- 1
  sim
}

# ---- complete weighted graph ------------------------------------------------
complete_graph_from_similarity <- function(similarity) {
  fns <- rownames(similarity)
  pairs <- as.data.frame(t(combn(fns, 2)), stringsAsFactors = FALSE)
  names(pairs) <- c("from", "to")
  edges <- pairs |>
    rowwise() |>
    mutate(similarity = similarity[from, to], distance = 1 - similarity) |>
    ungroup()
  igraph::graph_from_data_frame(edges, directed = FALSE, vertices = tibble(name = fns))
}

# ---- series → set of functions ----------------------------------------------
series_function_sets <- function(dataset, include_offmap) {
  dat <- dataset$semantic_counts
  if (!include_offmap) dat <- filter(dat, semantic_function %in% on_map_functions)
  dat |>
    distinct(series_key, semantic_function) |>
    group_by(series_key) |>
    summarise(functions = list(sort(unique(semantic_function))), .groups = "drop")
}

# ---- connectivity violations (holes) ----------------------------------------
connectivity_violations <- function(graph, series_sets) {
  verts <- igraph::V(graph)$name
  details <- series_sets |>
    mutate(
      available = map(functions, ~ intersect(.x, verts)),
      n_functions = map_int(available, length),
      n_components = map_int(available, function(fn_set) {
        if (length(fn_set) <= 1) return(1L)
        as.integer(igraph::components(igraph::induced_subgraph(graph, vids = fn_set))$no)
      }),
      violations = pmax(n_components - 1L, 0L)
    )
  list(
    total = sum(details$violations),
    violating_series = sum(details$violations > 0),
    details = details
  )
}

# ---- minimum spanning tree --------------------------------------------------
minimum_spanning_tree <- function(complete_graph) {
  tree <- igraph::mst(complete_graph, weights = igraph::E(complete_graph)$distance)
  igraph::E(tree)$edge_type <- "MST"
  tree
}

# ---- greedy edge addition to remove holes -----------------------------------
augment_graph <- function(graph, similarity, series_sets, max_edges = 100L) {
  cur_graph <- graph
  cur_score <- connectivity_violations(cur_graph, series_sets)$total
  history <- tibble(
    step = 0L, from = NA_character_, to = NA_character_, similarity = NA_real_,
    violations_before = cur_score, violations_after = cur_score, reduction = 0L
  )
  all_pairs <- as.data.frame(t(combn(igraph::V(cur_graph)$name, 2)), stringsAsFactors = FALSE)
  names(all_pairs) <- c("from", "to")
  
  for (step in seq_len(max_edges)) {
    if (cur_score == 0) break
    candidates <- all_pairs |>
      filter(!map2_lgl(from, to, ~ igraph::are_adjacent(cur_graph, .x, .y)))
    if (nrow(candidates) == 0) break
    evaluated <- candidates |>
      mutate(
        similarity = map2_dbl(from, to, ~ similarity[.x, .y]),
        score_after = map2_int(from, to, function(a, b) {
          tg <- igraph::add_edges(cur_graph, c(a, b))
          connectivity_violations(tg, series_sets)$total
        }),
        reduction = cur_score - score_after
      ) |>
      arrange(desc(reduction), desc(similarity), from, to)
    best <- evaluated[1, ]
    if (best$reduction[[1]] <= 0) break
    cur_graph <- igraph::add_edges(
      cur_graph, c(best$from[[1]], best$to[[1]]),
      attr = list(
        similarity = best$similarity[[1]],
        distance = 1 - best$similarity[[1]],
        edge_type = "added"
      )
    )
    history <- bind_rows(history, tibble(
      step = step, from = best$from[[1]], to = best$to[[1]],
      similarity = best$similarity[[1]], violations_before = cur_score,
      violations_after = best$score_after[[1]], reduction = best$reduction[[1]]
    ))
    cur_score <- best$score_after[[1]]
  }
  list(graph = cur_graph, history = history)
}

# ---- main analysis for one model --------------------------------------------
analyse_model <- function(dataset, dataset_name, model_name, use_contexts, include_offmap) {
  profile <- build_profile_matrix(dataset, use_contexts, include_offmap)
  similarity <- jaccard_matrix(profile)
  complete_g <- complete_graph_from_similarity(similarity)
  series_sets <- series_function_sets(dataset, include_offmap)
  
  mst_g <- minimum_spanning_tree(complete_g)
  before <- connectivity_violations(mst_g, series_sets)
  augmented <- augment_graph(mst_g, similarity, series_sets, max_added_edges)
  final_g <- augmented$graph
  after <- connectivity_violations(final_g, series_sets)
  
  # node sizes (number of series containing the function)
  node_freq <- dataset$semantic_counts |>
    filter(semantic_function %in% igraph::V(final_g)$name) |>
    distinct(series_key, semantic_function) |>
    count(semantic_function, name = "series_frequency")
  igraph::V(final_g)$series_frequency <- node_freq$series_frequency[
    match(igraph::V(final_g)$name, node_freq$semantic_function)
  ]
  igraph::V(final_g)$series_frequency[is.na(igraph::V(final_g)$series_frequency)] <- 0
  
  edge_table <- igraph::as_data_frame(final_g, what = "edges") |>
    mutate(dataset = dataset_name, model = model_name, .before = 1)
  
  similarity_long <- as.data.frame(as.table(similarity), stringsAsFactors = FALSE) |>
    as_tibble() |>
    rename(function_1 = Var1, function_2 = Var2, similarity = Freq) |>
    filter(function_1 < function_2) |>
    mutate(dataset = dataset_name, model = model_name, .before = 1)
  
  metrics <- tibble(
    dataset = dataset_name,
    model = model_name,
    use_contexts = use_contexts,
    include_offmap = include_offmap,
    n_functions = igraph::vcount(final_g),
    mst_edges = igraph::ecount(mst_g),
    added_edges = sum(igraph::E(final_g)$edge_type == "added"),
    total_edges = igraph::ecount(final_g),
    mst_violations = before$total,
    mst_violating_series = before$violating_series,
    final_violations = after$total,
    final_violating_series = after$violating_series,
    mean_edge_similarity = mean(igraph::E(final_g)$similarity, na.rm = TRUE),
    total_edge_distance = sum(igraph::E(final_g)$distance, na.rm = TRUE)
  )
  
  message(sprintf("  %s: functions=%d, MST edges=%d, added=%d, violations %d→%d",
                  dataset_name, metrics$n_functions, metrics$mst_edges,
                  metrics$added_edges, metrics$mst_violations, metrics$final_violations))
  
  list(
    profile = profile,
    similarity = similarity,
    graph = final_g,
    mst_graph = mst_g,
    series_sets = series_sets,
    metrics = metrics,
    edges = edge_table,
    similarities = similarity_long,
    history = augmented$history |> mutate(dataset = dataset_name, model = model_name, .before = 1),
    violations = after$details |> mutate(dataset = dataset_name, model = model_name, .before = 1)
  )
}

# ---- model definitions ------------------------------------------------------
models <- tibble::tribble(
  ~model,            ~use_contexts, ~include_offmap,
  "functions_onmap", FALSE,         FALSE,
  "functions_all",   FALSE,         TRUE,
  "contexts_onmap",  TRUE,          FALSE,
  "contexts_all",    TRUE,          TRUE
)

# ---- read data --------------------------------------------------------------
datasets <- list(
  "Map 1" = read_dataset(map1_file, "Map 1"),
  "Map 2" = read_dataset(map2_file, "Map 2")
)

# ---- run all models ---------------------------------------------------------
results <- list()
for (i in seq_len(nrow(models))) {
  current <- models[i, ]
  model_name <- current$model[[1]]
  message("\nAnalyzing model: ", model_name)
  for (ds_name in names(datasets)) {
    key <- paste(ds_name, model_name, sep = "__")
    results[[key]] <- analyse_model(
      dataset = datasets[[ds_name]],
      dataset_name = ds_name,
      model_name = model_name,
      use_contexts = current$use_contexts[[1]],
      include_offmap = current$include_offmap[[1]]
    )
  }
}

# ---- collect output tables --------------------------------------------------
metrics_all      <- map_dfr(results, "metrics") |> arrange(model, dataset)
edges_all        <- map_dfr(results, "edges") |> arrange(model, dataset, edge_type, desc(similarity))
similarities_all <- map_dfr(results, "similarities") |> arrange(model, dataset, desc(similarity))
history_all      <- map_dfr(results, "history") |> arrange(model, dataset, step)
violations_all   <- map_dfr(results, "violations") |>
  select(dataset, model, series_key, n_functions, n_components, violations) |>
  arrange(model, dataset, desc(violations), series_key)

# compare edges between Map 1 and Map 2
edge_comparison <- edges_all |>
  mutate(edge_key = map2_chr(from, to, ~ paste(sort(c(.x, .y)), collapse = " -- "))) |>
  select(dataset, model, edge_key, edge_type, similarity) |>
  pivot_wider(names_from = dataset, values_from = c(edge_type, similarity)) |>
  mutate(
    present_map1 = !is.na(`edge_type_Map 1`),
    present_map2 = !is.na(`edge_type_Map 2`),
    shared_edge  = present_map1 & present_map2
  ) |>
  arrange(model, desc(shared_edge), edge_key)

model_comparison <- metrics_all |>
  pivot_wider(
    names_from = dataset,
    values_from = c(n_functions, mst_edges, added_edges, total_edges,
                    mst_violations, mst_violating_series,
                    final_violations, final_violating_series,
                    mean_edge_similarity, total_edge_distance)
  )

# ---- save to Excel ----------------------------------------------------------
message("\nWriting Excel report: ", output_file)
if (file.exists(output_file)) file.remove(output_file)
wb <- openxlsx::createWorkbook()

add_sheet <- function(name, data) {
  openxlsx::addWorksheet(wb, name)
  openxlsx::writeData(wb, sheet = name, x = data, withFilter = TRUE)
  openxlsx::freezePane(wb, sheet = name, firstRow = TRUE)
  if (ncol(data) > 0)
    openxlsx::setColWidths(wb, sheet = name, cols = seq_len(ncol(data)), widths = "auto")
}

add_sheet("metrics", metrics_all)
add_sheet("model_comparison", model_comparison)
add_sheet("edges", edges_all)
add_sheet("edge_comparison", edge_comparison)
add_sheet("augmentation", history_all)
add_sheet("violations", violations_all)
add_sheet("similarities", similarities_all)

openxlsx::saveWorkbook(wb, output_file, overwrite = TRUE)
message("Excel file saved (", round(file.info(output_file)$size / 1024), " KB)")

# ---- console summary --------------------------------------------------------
message("\nAnalysis complete. Summary metrics:\n")
print(metrics_all, n = Inf, width = Inf)