1use rudb_common::{Error, Result};
25
26const SUPERBLOCK_BITS: usize = 4096;
28
29const BLOCK_BITS: usize = 512;
31
32const BLOCKS_PER_SUPERBLOCK: usize = SUPERBLOCK_BITS / BLOCK_BITS;
34
35const BLOCK_WORDS: usize = BLOCK_BITS / 64;
37
38const SAMPLE: u64 = 4096;
44
45#[derive(Debug, Clone)]
53pub(crate) struct Rank {
54 superblocks: Vec<u32>,
55 blocks: Vec<u16>,
56}
57
58impl Rank {
59 pub(crate) fn build(bits: &[u64]) -> Self {
60 let blocks = bits.len().div_ceil(BLOCK_WORDS);
61 let mut index = Self {
62 superblocks: Vec::with_capacity(blocks.div_ceil(BLOCKS_PER_SUPERBLOCK)),
63 blocks: Vec::with_capacity(blocks),
64 };
65 let mut total = 0_u32;
66 let mut within = 0_u16;
67 for block in 0..blocks {
68 if block % BLOCKS_PER_SUPERBLOCK == 0 {
69 index.superblocks.push(total);
70 within = 0;
71 }
72 index.blocks.push(within);
73 let words = block * BLOCK_WORDS;
74 let ones: u32 = bits[words..(words + BLOCK_WORDS).min(bits.len())]
75 .iter()
76 .map(|word| word.count_ones())
77 .sum();
78 total += ones;
79 #[expect(
82 clippy::cast_possible_truncation,
83 reason = "a superblock holds at most 4096 bits, which fits a u16"
84 )]
85 let ones = ones as u16;
86 within += ones;
87 }
88 index
89 }
90
91 pub(crate) fn rank(&self, bits: &[u64], at: usize) -> u64 {
93 let block = at / BLOCK_BITS;
94 let superblock = block / BLOCKS_PER_SUPERBLOCK;
95 let mut count = u64::from(self.superblocks[superblock]) + u64::from(self.blocks[block]);
96 let from = block * BLOCK_WORDS;
97 let word = at / 64;
98 for whole in &bits[from..word] {
99 count += u64::from(whole.count_ones());
100 }
101 let remainder = at % 64;
102 if remainder != 0 {
103 let mask = (1_u64 << remainder) - 1;
104 count += u64::from((bits[word] & mask).count_ones());
105 }
106 count
107 }
108
109 pub(crate) fn bytes(&self) -> usize {
110 self.superblocks.len() * size_of::<u32>() + self.blocks.len() * size_of::<u16>()
111 }
112
113 pub(crate) fn shape(words: usize) -> (usize, usize) {
118 let blocks = words.div_ceil(BLOCK_WORDS);
119 (blocks, blocks.div_ceil(BLOCKS_PER_SUPERBLOCK))
120 }
121
122 pub(crate) fn write(&self, out: &mut Vec<u8>) {
123 for count in &self.superblocks {
124 out.extend_from_slice(&count.to_le_bytes());
125 }
126 for offset in &self.blocks {
127 out.extend_from_slice(&offset.to_le_bytes());
128 }
129 }
130
131 pub(crate) fn read(bytes: &[u8], words: usize) -> Result<Self> {
133 let (blocks, superblocks) = Self::shape(words);
134 let split = superblocks * size_of::<u32>();
135 if bytes.len() != split + blocks * size_of::<u16>() {
136 return Err(malformed(
137 "a dense key map's rank index is not the size its range implies",
138 ));
139 }
140 Ok(Self {
141 superblocks: bytes[..split]
142 .chunks_exact(size_of::<u32>())
143 .map(|word| u32::from_le_bytes(word.try_into().expect("four bytes")))
144 .collect(),
145 blocks: bytes[split..]
146 .chunks_exact(size_of::<u16>())
147 .map(|word| u16::from_le_bytes(word.try_into().expect("two bytes")))
148 .collect(),
149 })
150 }
151
152 fn ones_before_superblock(&self, superblock: usize) -> u64 {
154 u64::from(self.superblocks[superblock])
155 }
156
157 fn ones_before_block(&self, block: usize) -> u64 {
159 u64::from(self.superblocks[block / BLOCKS_PER_SUPERBLOCK]) + u64::from(self.blocks[block])
160 }
161}
162
163#[derive(Debug, Clone)]
170pub struct BitVector {
171 words: Vec<u64>,
172 len: usize,
174 ones: u64,
175 rank: Rank,
176 ones_sample: Vec<u32>,
178 zeros_sample: Vec<u32>,
180}
181
182impl BitVector {
183 pub fn new(words: Vec<u64>, len: usize) -> Result<Self> {
192 if words.len() != len.div_ceil(64) {
193 return Err(malformed("a bit vector's words do not match its length"));
194 }
195 let tail = len % 64;
196 if tail != 0 && words[len / 64] >> tail != 0 {
197 return Err(malformed("a bit vector has bits set past its length"));
198 }
199 let rank = Rank::build(&words);
200 let ones = words.iter().map(|word| u64::from(word.count_ones())).sum();
201 let mut vector =
202 Self { words, len, ones, rank, ones_sample: Vec::new(), zeros_sample: Vec::new() };
203 vector.sample();
204 Ok(vector)
205 }
206
207 #[must_use]
209 pub fn len(&self) -> usize {
210 self.len
211 }
212
213 #[must_use]
215 pub fn is_empty(&self) -> bool {
216 self.len == 0
217 }
218
219 #[must_use]
221 pub fn ones(&self) -> u64 {
222 self.ones
223 }
224
225 #[must_use]
227 pub fn zeros(&self) -> u64 {
228 bits(self.len) - self.ones
229 }
230
231 #[must_use]
233 pub fn rank1(&self, at: usize) -> u64 {
234 if at >= self.len {
235 return self.ones;
236 }
237 self.rank.rank(&self.words, at)
238 }
239
240 #[must_use]
242 pub fn rank0(&self, at: usize) -> u64 {
243 let at = at.min(self.len);
244 bits(at) - self.rank1(at)
245 }
246
247 #[must_use]
249 pub fn select1(&self, nth: u64) -> Option<usize> {
250 if nth >= self.ones {
251 return None;
252 }
253 Some(self.select(nth, true))
254 }
255
256 #[must_use]
258 pub fn select0(&self, nth: u64) -> Option<usize> {
259 if nth >= self.zeros() {
260 return None;
261 }
262 Some(self.select(nth, false))
263 }
264
265 pub(crate) fn words(&self) -> &[u64] {
270 &self.words
271 }
272
273 #[must_use]
275 pub fn bytes(&self) -> usize {
276 self.words.len() * size_of::<u64>() + self.rank.bytes()
277 }
278
279 #[must_use]
283 pub fn bytes_for(len: usize) -> usize {
284 let words = len.div_ceil(64);
285 let (blocks, superblocks) = Rank::shape(words);
286 words * size_of::<u64>() + superblocks * size_of::<u32>() + blocks * size_of::<u16>()
287 }
288
289 pub fn write(&self, out: &mut Vec<u8>) {
291 for word in &self.words {
292 out.extend_from_slice(&word.to_le_bytes());
293 }
294 self.rank.write(out);
295 }
296
297 pub fn read(bytes: &[u8], len: usize) -> Result<Self> {
313 let words = len.div_ceil(64);
314 let bitmap = words * size_of::<u64>();
315 if bytes.len() < bitmap {
316 return Err(malformed("a bit vector is shorter than its length implies"));
317 }
318 let held = bytes[..bitmap]
319 .chunks_exact(size_of::<u64>())
320 .map(|word| u64::from_le_bytes(word.try_into().expect("eight bytes")))
321 .collect::<Vec<u64>>();
322 let rank = Rank::read(&bytes[bitmap..], words)?;
323 let tail = len % 64;
324 if tail != 0 && held[len / 64] >> tail != 0 {
325 return Err(malformed("a bit vector has bits set past its length"));
326 }
327 let ones = held.iter().map(|word| u64::from(word.count_ones())).sum();
328 let mut vector = Self {
329 words: held,
330 len,
331 ones,
332 rank,
333 ones_sample: Vec::new(),
334 zeros_sample: Vec::new(),
335 };
336 vector.sample();
337 Ok(vector)
338 }
339
340 fn sample(&mut self) {
348 self.ones_sample = self.samples(true, self.ones);
349 self.zeros_sample = self.samples(false, self.zeros());
350 }
351
352 fn samples(&self, set: bool, total: u64) -> Vec<u32> {
353 let superblocks = self.rank.superblocks.len();
354 let mut sample = Vec::with_capacity(usize::try_from(total.div_ceil(SAMPLE)).unwrap_or(0));
355 let mut at = 0_usize;
356 for group in 0..total.div_ceil(SAMPLE) {
357 let target = group * SAMPLE;
358 while at + 1 < superblocks && self.before(set, at + 1) <= target {
359 at += 1;
360 }
361 sample.push(u32::try_from(at).unwrap_or(u32::MAX));
362 }
363 sample
364 }
365
366 fn before(&self, set: bool, superblock: usize) -> u64 {
368 let ones = self.rank.ones_before_superblock(superblock);
369 if set { ones } else { bits(superblock * SUPERBLOCK_BITS) - ones }
370 }
371
372 fn select(&self, nth: u64, set: bool) -> usize {
379 if self.rank.superblocks.is_empty() {
380 return self.len;
381 }
382 let samples = if set { &self.ones_sample } else { &self.zeros_sample };
383 let last = self.rank.superblocks.len() - 1;
384 let group = usize::try_from(nth / SAMPLE).unwrap_or(usize::MAX);
385 let from = samples.get(group).map_or(0, |at| *at as usize);
386 let to = samples.get(group + 1).map_or(last, |at| *at as usize);
387 let (mut low, mut high) = (from, to);
388 while low < high {
389 let middle = low + (high - low).div_ceil(2);
392 if self.before(set, middle) <= nth {
393 low = middle;
394 } else {
395 high = middle - 1;
396 }
397 }
398 let superblock = low;
399 let blocks = self.rank.blocks.len();
400 let first = superblock * BLOCKS_PER_SUPERBLOCK;
401 let within = |block: usize| -> u64 {
402 let ones = self.rank.ones_before_block(block);
403 if set { ones } else { bits(block * BLOCK_BITS) - ones }
404 };
405 let mut block = first;
406 for candidate in first..(first + BLOCKS_PER_SUPERBLOCK).min(blocks) {
407 if within(candidate) <= nth {
408 block = candidate;
409 } else {
410 break;
411 }
412 }
413 let mut before = within(block);
414 for word in block * BLOCK_WORDS..self.words.len() {
415 let held = if set { self.words[word] } else { !self.words[word] };
416 let here = u64::from(held.count_ones());
417 if before + here > nth {
418 #[expect(
419 clippy::cast_possible_truncation,
420 reason = "a word holds at most 64 bits, so the offset within it fits a u32"
421 )]
422 let offset = (nth - before) as u32;
423 return word * 64 + nth_set(held, offset) as usize;
424 }
425 before += here;
426 }
427 self.len
431 }
432}
433
434fn bits(at: usize) -> u64 {
439 u64::try_from(at).unwrap_or(u64::MAX)
440}
441
442pub(crate) fn nth_set(mut word: u64, nth: u32) -> u32 {
448 for _ in 0..nth {
449 word &= word - 1;
450 }
451 word.trailing_zeros()
452}
453
454fn malformed(message: impl Into<String>) -> Error {
455 Error::invalid_input(format!("invalid rudb bit vector: {}", message.into()))
456}
457
458#[cfg(test)]
459mod tests {
460 use super::*;
461
462 fn vector(bits: &[bool]) -> BitVector {
464 let mut words = vec![0_u64; bits.len().div_ceil(64)];
465 for (at, bit) in bits.iter().enumerate() {
466 if *bit {
467 words[at / 64] |= 1 << (at % 64);
468 }
469 }
470 BitVector::new(words, bits.len()).expect("build")
471 }
472
473 fn agrees(bits: &[bool]) {
475 let built = vector(bits);
476 let (mut ones, mut zeros) = (Vec::new(), Vec::new());
477 for (at, bit) in bits.iter().enumerate() {
478 assert_eq!(built.rank1(at), ones.len() as u64, "rank1 at {at}");
479 assert_eq!(built.rank0(at), zeros.len() as u64, "rank0 at {at}");
480 if *bit { ones.push(at) } else { zeros.push(at) }
481 }
482 assert_eq!(built.ones(), ones.len() as u64);
483 assert_eq!(built.zeros(), zeros.len() as u64);
484 for (nth, at) in ones.iter().enumerate() {
485 assert_eq!(built.select1(nth as u64), Some(*at), "select1 of {nth}");
486 }
487 for (nth, at) in zeros.iter().enumerate() {
488 assert_eq!(built.select0(nth as u64), Some(*at), "select0 of {nth}");
489 }
490 assert_eq!(built.select1(ones.len() as u64), None, "there is no one past the last");
491 assert_eq!(built.select0(zeros.len() as u64), None, "there is no zero past the last");
492 }
493
494 #[test]
495 fn an_empty_vector_answers_nothing_rather_than_panicking() {
496 let built = vector(&[]);
497 assert!(built.is_empty());
498 assert_eq!(built.ones(), 0);
499 assert_eq!(built.zeros(), 0);
500 assert_eq!(built.select1(0), None);
501 assert_eq!(built.select0(0), None);
502 assert_eq!(built.rank1(0), 0);
503 assert_eq!(built.rank0(9), 0);
504 }
505
506 #[test]
507 fn a_vector_of_one_bit_each_way_agrees_with_counting() {
508 agrees(&[true]);
509 agrees(&[false]);
510 }
511
512 #[test]
513 fn alternating_bits_agree_with_counting_across_a_word_boundary() {
514 agrees(&(0..200).map(|at| at % 2 == 0).collect::<Vec<bool>>());
515 }
516
517 #[test]
518 fn a_vector_longer_than_a_superblock_agrees_with_counting() {
519 agrees(&(0..10_000).map(|at| at % 7 == 0).collect::<Vec<bool>>());
522 }
523
524 #[test]
525 fn a_vector_that_is_almost_all_ones_agrees_with_counting() {
526 agrees(&(0..20_000).map(|at| at % 1000 != 0).collect::<Vec<bool>>());
529 }
530
531 #[test]
532 fn a_vector_that_is_almost_all_zeros_agrees_with_counting() {
533 agrees(&(0..20_000).map(|at| at % 1000 == 0).collect::<Vec<bool>>());
534 }
535
536 #[test]
537 fn a_vector_of_all_ones_and_one_of_all_zeros_both_agree() {
538 agrees(&vec![true; 5000]);
539 agrees(&vec![false; 5000]);
540 }
541
542 #[test]
543 fn a_run_of_ones_longer_than_the_select_sample_is_found() {
544 let mut bits = vec![false; 3];
547 bits.extend(std::iter::repeat_n(true, 9000));
548 bits.push(false);
549 agrees(&bits);
550 }
551
552 #[test]
553 fn rank_past_the_end_saturates_rather_than_reading_past_it() {
554 let built = vector(&[true, false, true]);
555 assert_eq!(built.rank1(3), 2);
556 assert_eq!(built.rank1(9999), 2);
557 assert_eq!(built.rank0(9999), 1);
558 }
559
560 #[test]
561 fn a_vector_survives_being_written_and_read_back() {
562 let bits = (0..5000).map(|at| at % 3 == 0).collect::<Vec<bool>>();
563 let built = vector(&bits);
564 let mut bytes = Vec::new();
565 built.write(&mut bytes);
566 assert_eq!(bytes.len(), built.bytes(), "bytes() is what write() writes");
567 let read = BitVector::read(&bytes, bits.len()).expect("read");
568 assert_eq!(read.ones(), built.ones());
569 for nth in 0..read.ones() {
570 assert_eq!(read.select1(nth), built.select1(nth));
571 }
572 for nth in 0..read.zeros() {
573 assert_eq!(read.select0(nth), built.select0(nth));
574 }
575 }
576
577 #[test]
578 fn a_word_count_that_does_not_match_the_length_is_refused() {
579 assert!(BitVector::new(vec![0; 2], 64).is_err());
580 assert!(BitVector::new(vec![0; 1], 65).is_err());
581 }
582
583 #[test]
584 fn a_bit_set_past_the_length_is_refused_rather_than_counted() {
585 assert!(BitVector::new(vec![1 << 40], 8).is_err());
588 let mut bytes = Vec::new();
589 vector(&[true, false, true]).write(&mut bytes);
590 bytes[0] |= 1 << 4;
591 assert!(BitVector::read(&bytes, 3).is_err());
592 }
593
594 #[test]
595 fn a_truncated_vector_is_refused_rather_than_read_past() {
596 let mut bytes = Vec::new();
597 vector(&(0..5000).map(|at| at % 3 == 0).collect::<Vec<bool>>()).write(&mut bytes);
598 assert!(BitVector::read(&bytes[..bytes.len() - 1], 5000).is_err());
599 assert!(BitVector::read(&bytes[..4], 5000).is_err());
600 }
601
602 #[test]
603 fn the_rank_index_costs_about_an_eighth_of_the_bitmap() {
604 let built = vector(&(0..1_000_000).map(|at| at % 5 == 0).collect::<Vec<bool>>());
607 let bitmap = 1_000_000 / 8;
608 assert!(built.bytes() > bitmap, "{} is not more than {bitmap}", built.bytes());
609 assert!(
610 built.bytes() < bitmap * 6 / 5,
611 "{} is more than a fifth over {bitmap}",
612 built.bytes()
613 );
614 }
615}