use std::fmt::Write as _;
use crate::descriptor::{BuiltinTypeId, DynamicHasher, FormatSink, Tracer, TypeDescriptor};
use crate::repr_c_vec::ReprCVec;
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
pub struct BitIndex(usize);
impl BitIndex {
pub const MAX: i64 = u32::MAX as i64;
pub const MAX_WORDS: usize = (Self::MAX as usize) / 64 + 1;
#[must_use]
pub const fn new(value: i64) -> Option<BitIndex> {
if value < 0 || value > Self::MAX {
return None;
}
Some(BitIndex(value as usize))
}
#[inline]
const fn word_and_bit(self) -> (usize, usize) {
(self.0 / 64, self.0 % 64)
}
}
#[repr(C)]
pub struct BitSetPayload {
pub words: ReprCVec<u64>,
}
const _: () = assert!(std::mem::offset_of!(BitSetPayload, words) == 0);
const _: () = assert!(std::mem::size_of::<BitSetPayload>() == 24);
const _: () = assert!(std::mem::align_of::<BitSetPayload>() == 8);
#[cfg(not(feature = "std-vec-payload"))]
pub const INLINE_BITSET_SITE: crate::repr_c_vec::InlineSliceSite =
crate::repr_c_vec::InlineSliceSite::new(
BuiltinTypeId::BitSet,
std::mem::align_of::<BitSetPayload>(),
std::mem::offset_of!(BitSetPayload, words),
std::mem::size_of::<u64>(),
);
impl BitSetPayload {
pub(crate) fn contains(&self, i: BitIndex) -> bool {
let (word, bit) = i.word_and_bit();
match self.words.get(word) {
Some(w) => (w >> bit) & 1 == 1,
None => false,
}
}
pub(crate) fn insert(&mut self, i: BitIndex) {
let (word, bit) = i.word_and_bit();
if self.words.len() <= word {
self.words.resize(word + 1, 0);
}
debug_assert!(
self.words.len() <= BitIndex::MAX_WORDS,
"a BitSet's word count is bounded by BitIndex::MAX, and generated \
code reads that bound as licence to skip the range test"
);
self.words[word] |= 1u64 << bit;
}
pub(crate) fn remove(&mut self, i: BitIndex) {
let (word, bit) = i.word_and_bit();
if let Some(w) = self.words.get_mut(word) {
*w &= !(1u64 << bit);
}
}
pub(crate) fn count(&self) -> usize {
self.words.iter().map(|w| w.count_ones() as usize).sum()
}
pub(crate) fn members(&self) -> impl Iterator<Item = i64> + '_ {
self.words.iter().enumerate().flat_map(|(word_idx, &word)| {
let mut bits = word;
std::iter::from_fn(move || {
if bits == 0 {
return None;
}
let bit = bits.trailing_zeros() as usize;
bits &= bits - 1; Some((word_idx * 64 + bit) as i64)
})
})
}
}
unsafe fn bitset_trace(_payload: *mut u8, _tracer: &mut dyn Tracer) {
}
unsafe fn bitset_drop(payload: *mut u8) {
unsafe { std::ptr::drop_in_place(payload as *mut BitSetPayload) };
}
unsafe fn bitset_format(payload: *const u8, out: &mut FormatSink<'_>) {
let p = unsafe { &*(payload as *const BitSetPayload) };
let _ = out.write_str("{");
for (i, value) in p.members().enumerate() {
if i > 0 {
let _ = out.write_str(", ");
}
let _ = write!(out, "{value}");
}
let _ = out.write_str("}");
}
unsafe fn bitset_equals(a: *const u8, b: *const u8) -> bool {
let pa = unsafe { &*(a as *const BitSetPayload) };
let pb = unsafe { &*(b as *const BitSetPayload) };
let len = pa.words.len().max(pb.words.len());
for i in 0..len {
let wa = pa.words.get(i).copied().unwrap_or(0);
let wb = pb.words.get(i).copied().unwrap_or(0);
if wa != wb {
return false;
}
}
true
}
unsafe fn bitset_hash(payload: *const u8, hasher: &mut dyn DynamicHasher) {
let p = unsafe { &*(payload as *const BitSetPayload) };
let mut acc: u64 = 0;
for value in p.members() {
let mut h = crate::descriptor::StructHasher::new();
h.write_bytes(&(value as u64).to_le_bytes());
acc ^= h.finish();
}
hasher.write_bytes(&acc.to_le_bytes());
}
pub static BITSET: TypeDescriptor = TypeDescriptor::builtin::<BitSetPayload>(
BuiltinTypeId::BitSet,
"BitSet",
bitset_trace,
bitset_drop,
bitset_format,
Some(bitset_equals),
Some(bitset_hash),
None,
)
.with_owned_bytes(bitset_owned_bytes);
impl BitSetPayload {
#[must_use]
pub(crate) fn owned_bytes(&self) -> usize {
self.words.capacity() * std::mem::size_of::<u64>()
}
}
unsafe fn bitset_owned_bytes(payload: *const u8) -> usize {
let p = unsafe { &*(payload as *const BitSetPayload) };
p.owned_bytes()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bitset_descriptor_reports_capabilities() {
assert!(BITSET.is_equatable() && BITSET.is_hashable());
assert_eq!(BITSET.name, "BitSet");
}
fn bit(i: i64) -> BitIndex {
BitIndex::new(i).expect("in-range test bit")
}
#[test]
fn a_bitsets_members_come_out_ascending() {
let mut b = BitSetPayload {
words: ReprCVec::new(),
};
for i in [130, 5, 64, 0, 63] {
b.insert(bit(i));
}
assert_eq!(b.members().collect::<Vec<_>>(), vec![0, 5, 63, 64, 130]);
let mut rendered = String::new();
unsafe {
bitset_format(
(&b as *const BitSetPayload).cast::<u8>(),
&mut crate::FormatSink::display(&mut rendered),
)
};
assert_eq!(rendered, "{0, 5, 63, 64, 130}");
let empty = BitSetPayload {
words: ReprCVec::new(),
};
assert_eq!(empty.members().count(), 0);
let cleared = BitSetPayload {
words: ReprCVec::from_vec(vec![0, 0]),
};
assert_eq!(cleared.members().count(), 0);
}
#[test]
fn bitset_insert_contains_count() {
let mut b = BitSetPayload {
words: ReprCVec::new(),
};
b.insert(bit(0));
b.insert(bit(63));
b.insert(bit(64));
b.insert(bit(1000));
assert!(b.contains(bit(0)));
assert!(b.contains(bit(63)));
assert!(b.contains(bit(64)));
assert!(b.contains(bit(1000)));
assert!(!b.contains(bit(1)));
assert!(!b.contains(bit(65)));
assert_eq!(b.count(), 4);
}
#[test]
fn bitset_remove_clears_bit() {
let mut b = BitSetPayload {
words: ReprCVec::new(),
};
b.insert(bit(5));
assert!(b.contains(bit(5)));
b.remove(bit(5));
assert!(!b.contains(bit(5)));
b.remove(bit(999));
}
#[test]
fn the_word_probe_generated_code_emits_answers_contains() {
fn probe(words: &[u64], member: i64) -> bool {
let word = (member as u64) >> 6;
if word >= words.len() as u64 {
return false;
}
let w = words[word as usize];
(w >> ((member as u64) & 63)) & 1 == 1
}
let mut b = BitSetPayload {
words: ReprCVec::new(),
};
for i in [0, 1, 63, 64, 65, 127, 128, 1000, 4095, 4096] {
b.insert(bit(i));
}
let corners = [
i64::MIN,
i64::MIN + 1,
-4096,
-65,
-64,
-1,
0,
1,
63,
64,
BitIndex::MAX - 1,
BitIndex::MAX,
BitIndex::MAX + 1,
i64::MAX - 1,
i64::MAX,
];
for member in corners {
let expected = BitIndex::new(member).is_some_and(|i| b.contains(i));
assert_eq!(
probe(&b.words, member),
expected,
"the word probe and `contains` disagree at {member}"
);
}
for member in -64_i64..4200 {
let expected = BitIndex::new(member).is_some_and(|i| b.contains(i));
assert_eq!(probe(&b.words, member), expected, "at {member}");
}
assert_eq!(BitIndex::MAX_WORDS, 1 << 26);
assert!((BitIndex::MAX as u64) >> 6 < BitIndex::MAX_WORDS as u64);
assert!(((BitIndex::MAX + 1) as u64) >> 6 >= BitIndex::MAX_WORDS as u64);
assert!((-1_i64 as u64) >> 6 >= BitIndex::MAX_WORDS as u64);
assert!((i64::MIN as u64) >> 6 >= BitIndex::MAX_WORDS as u64);
}
#[test]
fn a_bit_outside_the_representable_range_has_no_index() {
assert!(BitIndex::new(-1).is_none(), "negative");
assert!(BitIndex::new(i64::MIN).is_none(), "most negative");
assert!(
BitIndex::new(BitIndex::MAX + 1).is_none(),
"one past the cap"
);
assert!(
BitIndex::new(i64::MAX).is_none(),
"the value that asked for a 10^16-word Vec"
);
assert!(BitIndex::new(0).is_some(), "zero is a member");
assert!(
BitIndex::new(BitIndex::MAX).is_some(),
"the cap is a member"
);
}
}