use std::sync::atomic::{AtomicUsize, Ordering};
const RANK_BITS: u32 = usize::BITS.ilog2();
const PARENT_BITS: u32 = usize::BITS - RANK_BITS;
pub const MAX_SIZE: usize = usize::MAX >> RANK_BITS;
pub struct UFRush {
nodes: Vec<AtomicUsize>,
}
impl UFRush {
pub fn new(size: usize) -> Self {
assert!(size <= MAX_SIZE);
Self {
nodes: (0..size).map(AtomicUsize::new).collect(),
}
}
pub fn size(&self) -> usize {
self.nodes.len()
}
pub fn same(&self, x: usize, y: usize) -> bool {
loop {
let x_rep = self.find(x);
let y_rep = self.find(y);
if x_rep == y_rep {
return true;
}
let x_node = self.nodes[x_rep].load(Ordering::Relaxed);
if x_rep == parent(x_node) {
return false;
}
}
}
pub fn find(&self, mut x: usize) -> usize {
assert!(x < self.size());
let mut x_node = self.nodes[x].load(Ordering::Relaxed);
while x != parent(x_node) {
let x_parent = parent(x_node);
let x_parent_node = self.nodes[x_parent].load(Ordering::Relaxed);
let x_parent_parent = parent(x_parent_node);
let x_new_node = encode(x_parent_parent, rank(x_node));
let _ = self.nodes[x].compare_exchange_weak(
x_node,
x_new_node,
Ordering::Release,
Ordering::Relaxed,
);
x = x_parent_parent;
x_node = self.nodes[x].load(Ordering::Relaxed);
}
x
}
pub fn unite(&self, x: usize, y: usize) -> bool {
loop {
let mut x_rep = self.find(x);
let mut y_rep = self.find(y);
if x_rep == y_rep {
return false;
}
let x_node = self.nodes[x_rep].load(Ordering::Relaxed);
let y_node = self.nodes[y_rep].load(Ordering::Relaxed);
let mut x_rank = rank(x_node);
let mut y_rank = rank(y_node);
if x_rank > y_rank || (x_rank == y_rank && x_rep > y_rep) {
std::mem::swap(&mut x_rep, &mut y_rep);
std::mem::swap(&mut x_rank, &mut y_rank);
}
let cur_value = encode(x_rep, x_rank);
let new_value = encode(y_rep, x_rank);
if self.nodes[x_rep]
.compare_exchange(cur_value, new_value, Ordering::Release, Ordering::Acquire)
.is_ok()
{
if x_rank == y_rank {
let cur_value = encode(y_rep, y_rank);
let new_value = encode(y_rep, y_rank + 1);
let _ = self.nodes[y_rep].compare_exchange_weak(
cur_value,
new_value,
Ordering::Release,
Ordering::Relaxed,
);
}
return true;
}
}
}
pub fn clear(&mut self) {
self.nodes
.iter_mut()
.enumerate()
.for_each(|(i, node)| node.store(i, Ordering::Relaxed));
}
}
unsafe impl Sync for UFRush {}
unsafe impl Send for UFRush {}
fn encode(parent: usize, rank: usize) -> usize {
parent | (rank << PARENT_BITS)
}
fn parent(n: usize) -> usize {
n & MAX_SIZE
}
fn rank(n: usize) -> usize {
n >> PARENT_BITS
}
#[cfg(test)]
mod tests {
use super::*;
use rand::prelude::*;
use std::collections::HashSet;
use std::sync::Arc;
use std::thread;
#[test]
fn test_new() {
let uf = UFRush::new(10);
assert_eq!(uf.size(), 10);
}
#[test]
fn test_find() {
let uf = UFRush::new(10);
assert_eq!(uf.find(5), 5);
}
#[test]
fn test_same() {
let uf = UFRush::new(10);
assert!(!uf.same(1, 2));
}
#[test]
fn test_unite() {
let uf = UFRush::new(10);
assert!(!uf.same(1, 2));
assert!(uf.unite(1, 2));
assert!(uf.same(1, 2));
}
#[test]
fn test_unite_already_united() {
let uf = UFRush::new(10);
assert!(uf.unite(1, 2));
assert!(!uf.unite(1, 2));
}
#[test]
fn test_clear() {
let mut uf = UFRush::new(10);
assert!(uf.unite(1, 2));
assert!(uf.same(1, 2));
uf.clear();
assert!(!uf.same(1, 2));
}
#[test]
fn test_multithreaded_build_cyclic_graph() {
let vertices = 100;
let uf = Arc::new(UFRush::new(vertices));
let handles: Vec<_> = (0..vertices)
.map(|n| {
let uf = Arc::clone(&uf);
thread::spawn(move || {
uf.unite(n, (n + 1) % vertices);
})
})
.collect();
for handle in handles {
handle.join().unwrap();
}
for n in 0..vertices - 1 {
assert!(uf.same(n, (n + 1) % vertices));
}
}
#[test]
fn test_multithreaded_cyclic_graph() {
assert!(is_cyclic(3, [(0, 1), (1, 2), (2, 0)]));
}
#[test]
fn test_multithreaded_acyclic_graph() {
assert!(!is_cyclic(4, [(0, 1), (1, 2), (2, 3)]));
}
#[test]
fn stress_test() {
let vertices = 100;
let mut edges = HashSet::with_capacity(5 * vertices);
let mut rng = rand::thread_rng();
edges.extend((0..vertices).map(|n| (n, (n + 1) % vertices)));
for _ in edges.len()..edges.capacity() {
let u = rng.gen_range(0..vertices);
let v = rng.gen_range(0..vertices);
if u != v {
edges.insert((u, v));
}
}
let mut edges: Vec<_> = edges.into_iter().collect();
for _ in 0..100 {
let uf = Arc::new(UFRush::new(vertices));
let handles: Vec<_> = edges
.iter()
.map(|&(u, v)| {
let uf = Arc::clone(&uf);
thread::spawn(move || uf.unite(u, v))
})
.collect();
let total_united = handles
.into_iter()
.map(|handle| handle.join().unwrap())
.filter(|&united| united)
.count();
assert_eq!(total_united, vertices - 1);
edges.shuffle(&mut rng);
}
}
fn is_cyclic<I>(vertices: usize, edges: I) -> bool
where
I: IntoIterator<Item = (usize, usize)>,
{
let uf = Arc::new(UFRush::new(vertices));
let handles: Vec<_> = edges
.into_iter()
.map(|(u, v)| {
let uf = Arc::clone(&uf);
thread::spawn(move || {
if uf.same(u, v) {
true
} else {
uf.unite(u, v);
false
}
})
})
.collect();
handles
.into_iter()
.map(|handle| handle.join().unwrap())
.any(|cyclic| cyclic)
}
}