use std::collections::HashMap;
pub fn stitch_components(
node_labels: Vec<(String, String)>,
boundary_edges: Vec<(String, String)>,
) -> Vec<(String, i64)> {
let mut uf = StringUnionFind::default();
for (name, local_root) in &node_labels {
uf.union(name, local_root);
}
for (a, b) in &boundary_edges {
uf.union(a, b);
}
let names = uf.names();
let mut root_of: HashMap<String, String> = HashMap::with_capacity(names.len());
for name in &names {
root_of.insert(name.clone(), uf.find(name));
}
let mut component_min: HashMap<String, String> = HashMap::new();
for name in &names {
let root = &root_of[name];
component_min
.entry(root.clone())
.and_modify(|m| {
if name < m {
*m = name.clone();
}
})
.or_insert_with(|| name.clone());
}
let mut mins: Vec<(&String, &String)> = component_min.iter().collect();
mins.sort_by(|a, b| a.1.cmp(b.1));
let mut id_of_root: HashMap<&String, i64> = HashMap::with_capacity(mins.len());
for (id, (root, _min)) in mins.iter().enumerate() {
id_of_root.insert(*root, id as i64);
}
let mut sorted_names = names;
sorted_names.sort();
sorted_names
.into_iter()
.map(|name| {
let root = &root_of[&name];
let id = id_of_root[root];
(name, id)
})
.collect()
}
#[derive(Default)]
struct StringUnionFind {
index: HashMap<String, usize>,
names: Vec<String>,
parent: Vec<usize>,
rank: Vec<u8>,
}
impl StringUnionFind {
fn intern(&mut self, name: &str) -> usize {
if let Some(&i) = self.index.get(name) {
return i;
}
let i = self.names.len();
self.index.insert(name.to_string(), i);
self.names.push(name.to_string());
self.parent.push(i);
self.rank.push(0);
i
}
fn find_idx(&mut self, mut x: usize) -> usize {
while self.parent[x] != x {
self.parent[x] = self.parent[self.parent[x]];
x = self.parent[x];
}
x
}
fn union(&mut self, a: &str, b: &str) {
let ia = self.intern(a);
let ib = self.intern(b);
let ra = self.find_idx(ia);
let rb = self.find_idx(ib);
if ra == rb {
return;
}
match self.rank[ra].cmp(&self.rank[rb]) {
std::cmp::Ordering::Less => self.parent[ra] = rb,
std::cmp::Ordering::Greater => self.parent[rb] = ra,
std::cmp::Ordering::Equal => {
self.parent[rb] = ra;
self.rank[ra] += 1;
}
}
}
fn find(&mut self, name: &str) -> String {
let i = self.index[name];
let r = self.find_idx(i);
self.names[r].clone()
}
fn names(&self) -> Vec<String> {
self.names.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn id_map(rows: &[(String, i64)]) -> HashMap<&str, i64> {
rows.iter().map(|(n, id)| (n.as_str(), *id)).collect()
}
#[test]
fn stitch_single_shard_two_components() {
let labels = vec![
("a".to_string(), "a".to_string()),
("b".to_string(), "a".to_string()),
("z".to_string(), "z".to_string()),
];
let rows = stitch_components(labels, Vec::new());
assert_eq!(rows.len(), 3);
let m = id_map(&rows);
assert_eq!(m["a"], m["b"]);
assert_ne!(m["a"], m["z"]);
assert_eq!(m["a"], 0);
assert_eq!(m["z"], 1);
}
#[test]
fn stitch_boundary_edge_merges_cross_shard_components() {
let labels = vec![
("a".to_string(), "a".to_string()),
("b".to_string(), "a".to_string()),
("c".to_string(), "c".to_string()),
("d".to_string(), "c".to_string()),
];
let boundary = vec![("b".to_string(), "c".to_string())];
let rows = stitch_components(labels, boundary);
assert_eq!(rows.len(), 4);
let m = id_map(&rows);
assert_eq!(m["a"], m["b"]);
assert_eq!(m["b"], m["c"]);
assert_eq!(m["c"], m["d"]);
assert_eq!(m["a"], 0);
}
#[test]
fn stitch_dense_ids_ordered_by_min_name() {
let labels = vec![
("z0".to_string(), "z0".to_string()),
("z1".to_string(), "z0".to_string()),
("a0".to_string(), "a0".to_string()),
("a1".to_string(), "a0".to_string()),
];
let rows = stitch_components(labels, Vec::new());
let m = id_map(&rows);
assert_eq!(m["a0"], 0, "component with min name a0 → id 0");
assert_eq!(m["a1"], 0);
assert_eq!(m["z0"], 1, "component with min name z0 → id 1");
assert_eq!(m["z1"], 1);
let order: Vec<&str> = rows.iter().map(|(n, _)| n.as_str()).collect();
assert_eq!(order, vec!["a0", "a1", "z0", "z1"]);
}
#[test]
fn stitch_boundary_edge_to_unreported_ghost() {
let labels = vec![("a".to_string(), "a".to_string())];
let boundary = vec![("a".to_string(), "g".to_string())];
let rows = stitch_components(labels, boundary);
let m = id_map(&rows);
assert_eq!(rows.len(), 2);
assert_eq!(m["a"], m["g"]);
}
#[test]
fn stitch_empty() {
let rows = stitch_components(Vec::new(), Vec::new());
assert!(rows.is_empty());
}
}