use alloc::vec::Vec;
pub(crate) type SetId = u32;
pub(crate) const NO_SET: SetId = SetId::MAX;
#[derive(Clone, Debug, Default)]
pub(crate) struct Sets {
width: usize,
words: Vec<u64>,
}
impl Sets {
pub(crate) fn new(bits: usize) -> Self {
Self {
width: bits.div_ceil(64).max(1),
words: Vec::new(),
}
}
pub(crate) fn alloc(&mut self) -> SetId {
let id = self.words.len() / self.width;
self.words.resize(self.words.len() + self.width, 0);
id as SetId
}
#[inline]
fn row(&self, id: SetId) -> &[u64] {
let start = id as usize * self.width;
&self.words[start..start + self.width]
}
#[inline]
fn row_mut(&mut self, id: SetId) -> &mut [u64] {
let start = id as usize * self.width;
&mut self.words[start..start + self.width]
}
#[inline]
pub(crate) fn contains(&self, id: SetId, bit: usize) -> bool {
let word = self.words[id as usize * self.width + (bit >> 6)];
(word >> (bit & 63)) & 1 != 0
}
pub(crate) fn insert(&mut self, id: SetId, bit: usize) -> bool {
let word = &mut self.row_mut(id)[bit >> 6];
let mask = 1u64 << (bit & 63);
let changed = *word & mask == 0;
*word |= mask;
changed
}
pub(crate) fn union(&mut self, dst: SetId, src: SetId) -> bool {
if dst == src {
return false;
}
let width = self.width;
let (d, s) = (dst as usize * width, src as usize * width);
let mut changed = false;
for i in 0..width {
let merged = self.words[d + i] | self.words[s + i];
changed |= merged != self.words[d + i];
self.words[d + i] = merged;
}
changed
}
pub(crate) fn join(&mut self, dst: SetId, src: SetId) -> bool {
dst != NO_SET && src != NO_SET && self.union(dst, src)
}
pub(crate) fn scratch(&self) -> Vec<u64> {
alloc::vec![0; self.width]
}
pub(crate) fn or_into(&self, buf: &mut [u64], id: SetId) {
for (word, add) in buf.iter_mut().zip(self.row(id)) {
*word |= add;
}
}
pub(crate) fn alloc_from(&mut self, buf: &[u64]) -> SetId {
let id = self.alloc();
self.row_mut(id).copy_from_slice(buf);
id
}
pub(crate) fn members(&self, id: SetId) -> impl Iterator<Item = usize> + '_ {
self.row(id).iter().enumerate().flat_map(|(i, &word)| {
let mut rest = word;
core::iter::from_fn(move || {
if rest == 0 {
return None;
}
let bit = rest.trailing_zeros() as usize;
rest &= rest - 1;
Some(i * 64 + bit)
})
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sets_insert_and_contains_across_words() {
let mut sets = Sets::new(130);
let a = sets.alloc();
assert_eq!(sets.members(a).count(), 0);
assert!(sets.insert(a, 0));
assert!(sets.insert(a, 64));
assert!(sets.insert(a, 129));
assert!(!sets.insert(a, 129));
assert!(sets.contains(a, 64));
assert!(!sets.contains(a, 63));
assert_eq!(sets.members(a).collect::<Vec<_>>(), [0, 64, 129]);
}
#[test]
fn test_sets_union_reports_change() {
let mut sets = Sets::new(10);
let a = sets.alloc();
let b = sets.alloc();
let _ = sets.insert(b, 3);
assert!(sets.union(a, b));
assert!(!sets.union(a, b));
assert!(!sets.union(a, a));
assert!(sets.contains(a, 3));
}
#[test]
fn test_sets_scratch_rows_and_join() {
let mut sets = Sets::new(70);
let a = sets.alloc();
let _ = sets.insert(a, 69);
let mut buf = sets.scratch();
sets.or_into(&mut buf, a);
let b = sets.alloc_from(&buf);
assert!(sets.contains(b, 69));
assert!(!sets.join(NO_SET, a));
assert!(!sets.join(a, NO_SET));
let c = sets.alloc();
assert!(sets.join(c, b));
}
#[test]
fn test_sets_are_independent() {
let mut sets = Sets::new(1);
let a = sets.alloc();
let b = sets.alloc();
let _ = sets.insert(a, 0);
assert!(!sets.contains(b, 0));
assert!(sets.contains(a, 0));
}
}