use crate::alloc_prelude::*;
#[derive(Clone, Default)]
#[cfg_attr(feature = "serde-serialize", derive(Serialize, Deserialize))]
pub(crate) struct UnionFind {
parents: Vec<u32>,
sizes: Vec<u32>,
}
impl UnionFind {
pub fn reset(&mut self, len: usize) {
self.parents.clear();
self.parents.extend(0..len as u32);
self.sizes.clear();
self.sizes.resize(len, 1);
}
pub fn find(&mut self, i: u32) -> u32 {
let mut i = i;
loop {
let p = self.parents[i as usize];
if p == i {
return i;
}
let gp = self.parents[p as usize];
self.parents[i as usize] = gp;
i = gp;
}
}
pub fn union(&mut self, a: u32, b: u32) {
let ra = self.find(a);
let rb = self.find(b);
if ra == rb {
return;
}
let (big, small) = if self.sizes[ra as usize] >= self.sizes[rb as usize] {
(ra, rb)
} else {
(rb, ra)
};
self.parents[small as usize] = big;
self.sizes[big as usize] += self.sizes[small as usize];
}
pub fn flatten(&mut self) {
for i in 0..self.parents.len() as u32 {
let root = self.find(i);
self.parents[i as usize] = root;
}
}
pub fn root(&self, i: u32) -> u32 {
let root = self.parents[i as usize];
debug_assert_eq!(self.parents[root as usize], root, "not flattened");
root
}
pub fn size(&self, root: u32) -> u32 {
self.sizes[root as usize]
}
}
#[cfg(test)]
mod test {
use super::UnionFind;
#[test]
fn union_find_components() {
let mut uf = UnionFind::default();
uf.reset(6);
uf.union(0, 1);
uf.union(2, 3);
uf.union(1, 2);
assert_eq!(uf.find(0), uf.find(3));
assert_ne!(uf.find(0), uf.find(4));
assert_ne!(uf.find(4), uf.find(5));
uf.flatten();
assert_eq!(uf.root(0), uf.root(3));
assert_ne!(uf.root(0), uf.root(4));
assert_eq!(uf.size(uf.root(0)), 4);
assert_eq!(uf.size(uf.root(4)), 1);
uf.reset(3);
assert_ne!(uf.find(0), uf.find(1));
}
}