use std::fmt;
use std::iter;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
#[derive(Clone, Default)]
pub struct SmallBitset {
near: u64,
#[allow(clippy::box_collection)]
far: Option<Box<Vec<u64>>>,
}
impl fmt::Debug for SmallBitset {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(f, "SmallBitset")?;
f.debug_list().entries(self.iter()).finish()
}
}
impl SmallBitset {
pub fn new() -> Self {
Self::default()
}
pub fn new_with_near(near: u64) -> Self {
Self { near, far: None }
}
pub fn insert(&mut self, val: usize) -> bool {
let (word, mask) = self.addr_mut(val);
let ret = 0 == (*word & mask);
*word |= mask;
ret
}
pub fn remove(&mut self, val: usize) -> bool {
let (word, mask) = self.addr_mut(val);
let ret = 0 != (*word & mask);
*word &= !mask;
ret
}
pub fn contains(&self, val: usize) -> bool {
let (word, mask) = self.addr(val);
0 != (word & mask)
}
pub fn iter(&self) -> impl Iterator<Item = usize> + '_ {
static EMPTY: Vec<u64> = Vec::new();
iter::once(self.near)
.chain(self.far.as_deref().unwrap_or(&EMPTY).iter().copied())
.enumerate()
.flat_map(move |(ix, word)| {
(0..64)
.filter(move |&bit| 0 != (word & (1 << bit)))
.map(move |bit| bit + ix * 64)
})
}
pub fn iter_far(&self) -> impl Iterator<Item = usize> + '_ {
self.iter().filter(|&i| i >= 64)
}
pub fn add_all(&mut self, rhs: &Self) {
self.near |= rhs.near;
match (&mut self.far, &rhs.far) {
(_, &None) => {},
(&mut Some(ref mut lfar), &Some(ref rfar)) => {
for (l, r) in lfar.iter_mut().zip(rfar.iter()) {
*l |= *r;
}
if rfar.len() > lfar.len() {
lfar.extend_from_slice(&rfar[lfar.len()..]);
}
},
(lfar, &Some(ref rfar)) => {
*lfar = Some(rfar.clone());
},
}
}
pub fn remove_all(&mut self, rhs: &Self) {
self.near &= !rhs.near;
if let (&mut Some(ref mut lfar), &Some(ref rfar)) =
(&mut self.far, &rhs.far)
{
for (l, r) in lfar.iter_mut().zip(rfar.iter()) {
*l &= !*r;
}
}
}
pub fn remove_complement(&mut self, rhs: &Self) {
self.near &= rhs.near;
match (&mut self.far, &rhs.far) {
(&mut None, _) => {},
(lfar, &None) => *lfar = None,
(&mut Some(ref mut lfar), &Some(ref rfar)) => {
if lfar.len() > rfar.len() {
lfar.truncate(rfar.len());
}
for (l, r) in lfar.iter_mut().zip(rfar.iter()) {
*l &= *r;
}
},
}
}
pub fn near_bits(&self) -> u64 {
self.near
}
pub fn has_far(&self) -> bool {
self.far.is_some()
}
fn addr_mut(&mut self, val: usize) -> (&mut u64, u64) {
if val < 64 {
(&mut self.near, 1 << val)
} else {
let ix = val / 64 - 1;
let far = self.far.get_or_insert_with(Box::default);
if far.len() <= ix {
far.resize(ix + 1, 0);
}
(&mut far[ix], 1 << (val % 64))
}
}
fn addr(&self, val: usize) -> (u64, u64) {
if val < 64 {
(self.near, 1 << val)
} else if let Some(far) = self.far.as_ref() {
let ix = val / 64 - 1;
(far.get(ix).copied().unwrap_or(0), 1 << (val % 64))
} else {
(0, 1 << (val % 64))
}
}
}
impl Serialize for SmallBitset {
fn serialize<S: Serializer>(
&self,
serializer: S,
) -> Result<S::Ok, S::Error> {
match self.far {
None => {
let near_array = [self.near];
if self.near == 0 {
&[] as &[u64]
} else {
&near_array as &[u64]
}
.serialize(serializer)
},
Some(ref far) => {
let mut elements = Vec::clone(far);
elements.push(self.near);
elements.serialize(serializer)
},
}
}
}
impl<'de> Deserialize<'de> for SmallBitset {
fn deserialize<D: Deserializer<'de>>(
deserializer: D,
) -> Result<Self, D::Error> {
let mut elements: Vec<u64> = Vec::deserialize(deserializer)?;
let near = elements.pop().unwrap_or(0);
Ok(SmallBitset {
near,
far: if elements.is_empty() {
None
} else {
Some(Box::new(elements))
},
})
}
}
impl std::iter::FromIterator<usize> for SmallBitset {
fn from_iter<T: IntoIterator<Item = usize>>(iter: T) -> Self {
let mut this = Self::new();
for bit in iter {
this.insert(bit);
}
this
}
}
#[cfg(test)]
impl From<Vec<usize>> for SmallBitset {
fn from(v: Vec<usize>) -> Self {
v.into_iter().collect()
}
}
impl std::cmp::PartialEq for SmallBitset {
fn eq(&self, rhs: &Self) -> bool {
if self.near != rhs.near {
return false;
}
match (&self.far, &rhs.far) {
(&None, &None) => true,
(&Some(ref v), &None) | (&None, &Some(ref v)) => {
v.iter().all(|&w| 0 == w)
},
(&Some(ref lhs), &Some(ref rhs)) => {
let len = lhs.len().max(rhs.len());
lhs.iter()
.copied()
.chain(std::iter::repeat(0))
.zip(rhs.iter().copied().chain(std::iter::repeat(0)))
.take(len)
.all(|(lhs, rhs)| lhs == rhs)
},
}
}
}
impl std::cmp::Eq for SmallBitset {}
#[cfg(test)]
mod test {
use proptest::prelude::*;
use serde_cbor;
use super::*;
#[test]
fn basic_operations() {
let mut bs = SmallBitset::new();
assert!(!bs.contains(0));
assert!(!bs.contains(100));
assert!(!bs.contains(usize::MAX));
assert!(bs.insert(0));
assert!(bs.insert(42));
assert!(!bs.insert(42));
assert!(bs.contains(0));
assert!(!bs.contains(1));
assert!(bs.contains(42));
assert_eq!(vec![0, 42], bs.iter().collect::<Vec<_>>());
assert!(bs.remove(0));
assert!(!bs.remove(0));
assert!(!bs.contains(0));
assert!(bs.contains(42));
assert_eq!(vec![42], bs.iter().collect::<Vec<_>>());
assert!(bs.insert(100));
assert!(bs.contains(100));
assert_eq!(vec![42, 100], bs.iter().collect::<Vec<_>>());
assert!(bs.insert(1000));
assert!(bs.contains(1000));
assert_eq!(vec![42, 100, 1000], bs.iter().collect::<Vec<_>>());
assert!(bs.remove(100));
assert!(!bs.contains(100));
assert_eq!(vec![42, 1000], bs.iter().collect::<Vec<_>>());
}
fn serde_flip(bs: &SmallBitset) {
let as_bytes = serde_cbor::to_vec(bs).unwrap();
let reread: SmallBitset =
serde_cbor::from_reader(&as_bytes[..]).unwrap();
assert_eq!(
bs.iter().collect::<Vec<_>>(),
reread.iter().collect::<Vec<_>>()
);
}
#[test]
fn test_serde() {
let mut bs = SmallBitset::new();
serde_flip(&bs);
bs.insert(0);
serde_flip(&bs);
bs.insert(42);
serde_flip(&bs);
bs.remove(0);
serde_flip(&bs);
bs.insert(100);
serde_flip(&bs);
bs.insert(1000);
serde_flip(&bs);
bs.remove(42);
serde_flip(&bs);
bs.remove(1000);
serde_flip(&bs);
}
#[test]
fn partial_eq() {
let mut bs1 = SmallBitset::new();
let mut bs2 = SmallBitset::new();
assert_eq!(bs1, bs2);
bs1.insert(1);
assert_ne!(bs1, bs2);
assert_ne!(bs2, bs1);
bs2.insert(1);
assert_eq!(bs1, bs2);
assert_eq!(bs2, bs1);
bs1.insert(99);
assert_ne!(bs1, bs2);
assert_ne!(bs2, bs1);
bs1.remove(99);
assert_eq!(bs1, bs2);
assert_eq!(bs2, bs1);
bs2.insert(999);
assert_ne!(bs1, bs2);
assert_ne!(bs2, bs1);
bs2.remove(999);
assert_eq!(bs1, bs2);
assert_eq!(bs2, bs1);
}
proptest! {
#[test]
fn add_all(
hs1 in prop::collection::hash_set(
0usize..128,
0..10,
),
hs2 in prop::collection::hash_set(
0usize..128,
0..10,
),
) {
let bs1 = hs1.iter().copied().collect::<SmallBitset>();
let bs2 = hs2.iter().copied().collect::<SmallBitset>();
let combined = hs1.iter().copied()
.chain(hs2.iter().copied())
.collect::<SmallBitset>();
let mut result = bs1.clone();
result.add_all(&bs2);
assert_eq!(combined, result);
result = bs2.clone();
result.add_all(&bs1);
assert_eq!(combined, result);
}
#[test]
fn remove_all(
hs1 in prop::collection::hash_set(
0usize..128,
0..10,
),
hs2 in prop::collection::hash_set(
0usize..128,
0..10,
),
) {
let bs1 = hs1.iter().copied().collect::<SmallBitset>();
let bs2 = hs2.iter().copied().collect::<SmallBitset>();
let bs1_minus_bs2 = hs1.iter().copied()
.filter(|e| !hs2.contains(e))
.collect::<SmallBitset>();
let bs2_minus_bs1 = hs2.iter().copied()
.filter(|e| !hs1.contains(e))
.collect::<SmallBitset>();
let mut result = bs1.clone();
result.remove_all(&bs2);
assert_eq!(bs1_minus_bs2, result);
result = bs2.clone();
result.remove_all(&bs1);
assert_eq!(bs2_minus_bs1, result);
}
#[test]
fn remove_complement(
hs1 in prop::collection::hash_set(
0usize..128,
0..10,
),
hs2 in prop::collection::hash_set(
0usize..128,
0..10,
),
) {
let bs1 = hs1.iter().copied().collect::<SmallBitset>();
let bs2 = hs2.iter().copied().collect::<SmallBitset>();
let bs1_intersect_bs2 = hs1.iter().copied()
.filter(|e| hs2.contains(e))
.collect::<SmallBitset>();
let mut result = bs1.clone();
result.remove_complement(&bs2);
assert_eq!(bs1_intersect_bs2, result);
result = bs2.clone();
result.remove_complement(&bs1);
assert_eq!(bs1_intersect_bs2, result);
}
}
}