use std::collections::HashSet;
#[derive(Clone, Debug)]
pub enum VisibleSet {
Sparse(HashSet<u32>),
Dense { words: Box<[u64]>, len: usize },
}
impl VisibleSet {
pub fn from_ids(ids: impl IntoIterator<Item = u32>) -> VisibleSet {
let ids: Vec<u32> = ids.into_iter().collect();
let Some(&max) = ids.iter().max() else {
return VisibleSet::Sparse(HashSet::new());
};
let span = (max as usize).saturating_add(1);
if ids.len().saturating_mul(64) < span {
return VisibleSet::Sparse(ids.into_iter().collect());
}
let mut words = vec![0u64; span.div_ceil(64)];
let mut len = 0usize;
for id in ids {
let bit = 1u64 << (id % 64);
let word = &mut words[id as usize / 64];
if *word & bit == 0 {
*word |= bit;
len += 1;
}
}
VisibleSet::from_words(words, len)
}
fn from_words(mut words: Vec<u64>, len: usize) -> VisibleSet {
while words.last() == Some(&0) {
words.pop();
}
let span = match words.last() {
None => return VisibleSet::Sparse(HashSet::new()),
Some(&top) => (words.len() - 1) * 64 + (64 - top.leading_zeros() as usize),
};
if len.saturating_mul(64) >= span {
VisibleSet::Dense {
words: words.into_boxed_slice(),
len,
}
} else {
VisibleSet::Sparse(bits(&words).collect())
}
}
#[inline]
pub fn contains(&self, id: u32) -> bool {
match self {
VisibleSet::Sparse(set) => set.contains(&id),
VisibleSet::Dense { words, .. } => words
.get(id as usize / 64)
.is_some_and(|w| w >> (id % 64) & 1 == 1),
}
}
pub fn len(&self) -> usize {
match self {
VisibleSet::Sparse(set) => set.len(),
VisibleSet::Dense { len, .. } => *len,
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn iter(&self) -> impl Iterator<Item = u32> + '_ {
let sparse = match self {
VisibleSet::Sparse(set) => Some(set.iter().copied()),
VisibleSet::Dense { .. } => None,
};
let dense = match self {
VisibleSet::Sparse(_) => None,
VisibleSet::Dense { words, .. } => Some(bits(words)),
};
sparse
.into_iter()
.flatten()
.chain(dense.into_iter().flatten())
}
pub fn intersect(&self, other: &VisibleSet) -> VisibleSet {
if let (VisibleSet::Dense { words: a, .. }, VisibleSet::Dense { words: b, .. }) =
(self, other)
{
let n = a.len().min(b.len());
let words: Vec<u64> = (0..n).map(|i| a[i] & b[i]).collect();
let len = words.iter().map(|w| w.count_ones() as usize).sum();
return VisibleSet::from_words(words, len);
}
let (small, large) = if self.len() <= other.len() {
(self, other)
} else {
(other, self)
};
VisibleSet::from_ids(small.iter().filter(|&id| large.contains(id)))
}
}
fn bits(words: &[u64]) -> impl Iterator<Item = u32> + '_ {
words.iter().enumerate().flat_map(|(w, &word)| {
(0..64u32)
.filter(move |b| word >> b & 1 == 1)
.map(move |b| (w * 64) as u32 + b)
})
}
impl FromIterator<u32> for VisibleSet {
fn from_iter<I: IntoIterator<Item = u32>>(ids: I) -> VisibleSet {
VisibleSet::from_ids(ids)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn variant(v: &VisibleSet) -> &'static str {
match v {
VisibleSet::Sparse(_) => "sparse",
VisibleSet::Dense { .. } => "dense",
}
}
#[test]
fn the_rule_picks_the_variant_and_never_the_answer() {
let at: VisibleSet = (0..100).map(|i| i * 64 + 63).collect();
assert_eq!(at.len(), 100);
assert_eq!(variant(&at), "dense", "len * 64 == span is dense");
let over: VisibleSet = (0..99)
.map(|i| i * 64 + 63)
.chain(std::iter::once(6463))
.collect();
assert_eq!(over.len(), 100);
assert_eq!(variant(&over), "sparse", "len * 64 < span is sparse");
for id in 0..7000u32 {
assert_eq!(at.contains(id), id % 64 == 63 && id < 6400, "at id {id}");
}
assert!(over.contains(6463));
assert!(!over.contains(6462));
}
#[test]
fn a_ten_id_mask_in_a_ten_million_id_space_stays_sparse() {
let far: VisibleSet = (0..10).map(|i| 9_999_990 + i).collect();
assert_eq!(variant(&far), "sparse");
assert_eq!(far.len(), 10);
assert!(far.contains(9_999_999));
assert!(!far.contains(9_999_989));
}
#[test]
fn an_empty_set_is_sparse_and_contains_nothing() {
let empty = VisibleSet::from_ids(std::iter::empty());
assert_eq!(variant(&empty), "sparse");
assert!(empty.is_empty());
assert!(!empty.contains(0));
}
#[test]
fn duplicates_do_not_inflate_the_count_the_rule_sees() {
let dupes: VisibleSet = std::iter::repeat_n(1_000_000u32, 50_000).collect();
assert_eq!(dupes.len(), 1);
assert_eq!(variant(&dupes), "sparse");
assert!(dupes.contains(1_000_000));
}
#[test]
fn iter_returns_exactly_the_members_of_either_variant() {
for ids in [vec![0u32, 1, 2, 63, 64, 65], vec![0u32, 9_999_999]] {
let v: VisibleSet = ids.iter().copied().collect();
let mut got: Vec<u32> = v.iter().collect();
got.sort_unstable();
assert_eq!(got, ids, "{} lost a member", variant(&v));
}
}
#[test]
fn intersect_narrows_on_every_pairing_of_variants() {
let dense_a: VisibleSet = (0..200u32).collect();
let dense_b: VisibleSet = (100..300u32).collect();
let sparse_a: VisibleSet = [5u32, 150, 9_999_999].into_iter().collect();
assert_eq!(variant(&dense_a), "dense");
assert_eq!(variant(&sparse_a), "sparse");
let dd = dense_a.intersect(&dense_b);
let mut got: Vec<u32> = dd.iter().collect();
got.sort_unstable();
assert_eq!(got, (100..200).collect::<Vec<u32>>());
for (l, r) in [(&dense_a, &sparse_a), (&sparse_a, &dense_a)] {
let out = l.intersect(r);
let mut got: Vec<u32> = out.iter().collect();
got.sort_unstable();
assert_eq!(got, vec![5, 150], "intersection is order-independent");
}
}
#[test]
fn intersect_of_two_dense_sets_respects_the_shorter_span() {
let short: VisibleSet = (0..64u32).collect();
let long: VisibleSet = (0..640u32).collect();
let out = short.intersect(&long);
assert_eq!(out.len(), 64);
assert!(out.contains(63));
assert!(!out.contains(64));
}
#[test]
fn a_dense_intersection_that_narrows_hard_demotes_itself() {
let a: VisibleSet = (0..200_000u32).filter(|i| i % 2 == 0).collect();
let b: VisibleSet = (0..3_200u32).chain(std::iter::once(199_998)).collect();
assert_eq!(variant(&a), "dense");
assert_eq!(
variant(&b),
"dense",
"both sides must take the word-wise arm"
);
let out = a.intersect(&b);
assert_eq!(out.len(), 1_601);
assert_eq!(variant(&out), "sparse");
assert!(out.contains(0) && out.contains(3_198) && out.contains(199_998));
assert!(!out.contains(1) && !out.contains(3_200));
}
}