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 pub fn write(&self, out: &mut Vec<u8>) {
281 for word in &self.words {
282 out.extend_from_slice(&word.to_le_bytes());
283 }
284 self.rank.write(out);
285 }
286
287 pub fn read(bytes: &[u8], len: usize) -> Result<Self> {
303 let words = len.div_ceil(64);
304 let bitmap = words * size_of::<u64>();
305 if bytes.len() < bitmap {
306 return Err(malformed("a bit vector is shorter than its length implies"));
307 }
308 let held = bytes[..bitmap]
309 .chunks_exact(size_of::<u64>())
310 .map(|word| u64::from_le_bytes(word.try_into().expect("eight bytes")))
311 .collect::<Vec<u64>>();
312 let rank = Rank::read(&bytes[bitmap..], words)?;
313 let tail = len % 64;
314 if tail != 0 && held[len / 64] >> tail != 0 {
315 return Err(malformed("a bit vector has bits set past its length"));
316 }
317 let ones = held.iter().map(|word| u64::from(word.count_ones())).sum();
318 let mut vector = Self {
319 words: held,
320 len,
321 ones,
322 rank,
323 ones_sample: Vec::new(),
324 zeros_sample: Vec::new(),
325 };
326 vector.sample();
327 Ok(vector)
328 }
329
330 fn sample(&mut self) {
338 self.ones_sample = self.samples(true, self.ones);
339 self.zeros_sample = self.samples(false, self.zeros());
340 }
341
342 fn samples(&self, set: bool, total: u64) -> Vec<u32> {
343 let superblocks = self.rank.superblocks.len();
344 let mut sample = Vec::with_capacity(usize::try_from(total.div_ceil(SAMPLE)).unwrap_or(0));
345 let mut at = 0_usize;
346 for group in 0..total.div_ceil(SAMPLE) {
347 let target = group * SAMPLE;
348 while at + 1 < superblocks && self.before(set, at + 1) <= target {
349 at += 1;
350 }
351 sample.push(u32::try_from(at).unwrap_or(u32::MAX));
352 }
353 sample
354 }
355
356 fn before(&self, set: bool, superblock: usize) -> u64 {
358 let ones = self.rank.ones_before_superblock(superblock);
359 if set { ones } else { bits(superblock * SUPERBLOCK_BITS) - ones }
360 }
361
362 fn select(&self, nth: u64, set: bool) -> usize {
369 if self.rank.superblocks.is_empty() {
370 return self.len;
371 }
372 let samples = if set { &self.ones_sample } else { &self.zeros_sample };
373 let last = self.rank.superblocks.len() - 1;
374 let group = usize::try_from(nth / SAMPLE).unwrap_or(usize::MAX);
375 let from = samples.get(group).map_or(0, |at| *at as usize);
376 let to = samples.get(group + 1).map_or(last, |at| *at as usize);
377 let (mut low, mut high) = (from, to);
378 while low < high {
379 let middle = low + (high - low).div_ceil(2);
382 if self.before(set, middle) <= nth {
383 low = middle;
384 } else {
385 high = middle - 1;
386 }
387 }
388 let superblock = low;
389 let blocks = self.rank.blocks.len();
390 let first = superblock * BLOCKS_PER_SUPERBLOCK;
391 let within = |block: usize| -> u64 {
392 let ones = self.rank.ones_before_block(block);
393 if set { ones } else { bits(block * BLOCK_BITS) - ones }
394 };
395 let mut block = first;
396 for candidate in first..(first + BLOCKS_PER_SUPERBLOCK).min(blocks) {
397 if within(candidate) <= nth {
398 block = candidate;
399 } else {
400 break;
401 }
402 }
403 let mut before = within(block);
404 for word in block * BLOCK_WORDS..self.words.len() {
405 let held = if set { self.words[word] } else { !self.words[word] };
406 let here = u64::from(held.count_ones());
407 if before + here > nth {
408 #[expect(
409 clippy::cast_possible_truncation,
410 reason = "a word holds at most 64 bits, so the offset within it fits a u32"
411 )]
412 let offset = (nth - before) as u32;
413 return word * 64 + nth_set(held, offset) as usize;
414 }
415 before += here;
416 }
417 self.len
421 }
422}
423
424fn bits(at: usize) -> u64 {
429 u64::try_from(at).unwrap_or(u64::MAX)
430}
431
432fn nth_set(mut word: u64, nth: u32) -> u32 {
438 for _ in 0..nth {
439 word &= word - 1;
440 }
441 word.trailing_zeros()
442}
443
444fn malformed(message: impl Into<String>) -> Error {
445 Error::invalid_input(format!("invalid rudb bit vector: {}", message.into()))
446}
447
448#[cfg(test)]
449mod tests {
450 use super::*;
451
452 fn vector(bits: &[bool]) -> BitVector {
454 let mut words = vec![0_u64; bits.len().div_ceil(64)];
455 for (at, bit) in bits.iter().enumerate() {
456 if *bit {
457 words[at / 64] |= 1 << (at % 64);
458 }
459 }
460 BitVector::new(words, bits.len()).expect("build")
461 }
462
463 fn agrees(bits: &[bool]) {
465 let built = vector(bits);
466 let (mut ones, mut zeros) = (Vec::new(), Vec::new());
467 for (at, bit) in bits.iter().enumerate() {
468 assert_eq!(built.rank1(at), ones.len() as u64, "rank1 at {at}");
469 assert_eq!(built.rank0(at), zeros.len() as u64, "rank0 at {at}");
470 if *bit { ones.push(at) } else { zeros.push(at) }
471 }
472 assert_eq!(built.ones(), ones.len() as u64);
473 assert_eq!(built.zeros(), zeros.len() as u64);
474 for (nth, at) in ones.iter().enumerate() {
475 assert_eq!(built.select1(nth as u64), Some(*at), "select1 of {nth}");
476 }
477 for (nth, at) in zeros.iter().enumerate() {
478 assert_eq!(built.select0(nth as u64), Some(*at), "select0 of {nth}");
479 }
480 assert_eq!(built.select1(ones.len() as u64), None, "there is no one past the last");
481 assert_eq!(built.select0(zeros.len() as u64), None, "there is no zero past the last");
482 }
483
484 #[test]
485 fn an_empty_vector_answers_nothing_rather_than_panicking() {
486 let built = vector(&[]);
487 assert!(built.is_empty());
488 assert_eq!(built.ones(), 0);
489 assert_eq!(built.zeros(), 0);
490 assert_eq!(built.select1(0), None);
491 assert_eq!(built.select0(0), None);
492 assert_eq!(built.rank1(0), 0);
493 assert_eq!(built.rank0(9), 0);
494 }
495
496 #[test]
497 fn a_vector_of_one_bit_each_way_agrees_with_counting() {
498 agrees(&[true]);
499 agrees(&[false]);
500 }
501
502 #[test]
503 fn alternating_bits_agree_with_counting_across_a_word_boundary() {
504 agrees(&(0..200).map(|at| at % 2 == 0).collect::<Vec<bool>>());
505 }
506
507 #[test]
508 fn a_vector_longer_than_a_superblock_agrees_with_counting() {
509 agrees(&(0..10_000).map(|at| at % 7 == 0).collect::<Vec<bool>>());
512 }
513
514 #[test]
515 fn a_vector_that_is_almost_all_ones_agrees_with_counting() {
516 agrees(&(0..20_000).map(|at| at % 1000 != 0).collect::<Vec<bool>>());
519 }
520
521 #[test]
522 fn a_vector_that_is_almost_all_zeros_agrees_with_counting() {
523 agrees(&(0..20_000).map(|at| at % 1000 == 0).collect::<Vec<bool>>());
524 }
525
526 #[test]
527 fn a_vector_of_all_ones_and_one_of_all_zeros_both_agree() {
528 agrees(&vec![true; 5000]);
529 agrees(&vec![false; 5000]);
530 }
531
532 #[test]
533 fn a_run_of_ones_longer_than_the_select_sample_is_found() {
534 let mut bits = vec![false; 3];
537 bits.extend(std::iter::repeat_n(true, 9000));
538 bits.push(false);
539 agrees(&bits);
540 }
541
542 #[test]
543 fn rank_past_the_end_saturates_rather_than_reading_past_it() {
544 let built = vector(&[true, false, true]);
545 assert_eq!(built.rank1(3), 2);
546 assert_eq!(built.rank1(9999), 2);
547 assert_eq!(built.rank0(9999), 1);
548 }
549
550 #[test]
551 fn a_vector_survives_being_written_and_read_back() {
552 let bits = (0..5000).map(|at| at % 3 == 0).collect::<Vec<bool>>();
553 let built = vector(&bits);
554 let mut bytes = Vec::new();
555 built.write(&mut bytes);
556 assert_eq!(bytes.len(), built.bytes(), "bytes() is what write() writes");
557 let read = BitVector::read(&bytes, bits.len()).expect("read");
558 assert_eq!(read.ones(), built.ones());
559 for nth in 0..read.ones() {
560 assert_eq!(read.select1(nth), built.select1(nth));
561 }
562 for nth in 0..read.zeros() {
563 assert_eq!(read.select0(nth), built.select0(nth));
564 }
565 }
566
567 #[test]
568 fn a_word_count_that_does_not_match_the_length_is_refused() {
569 assert!(BitVector::new(vec![0; 2], 64).is_err());
570 assert!(BitVector::new(vec![0; 1], 65).is_err());
571 }
572
573 #[test]
574 fn a_bit_set_past_the_length_is_refused_rather_than_counted() {
575 assert!(BitVector::new(vec![1 << 40], 8).is_err());
578 let mut bytes = Vec::new();
579 vector(&[true, false, true]).write(&mut bytes);
580 bytes[0] |= 1 << 4;
581 assert!(BitVector::read(&bytes, 3).is_err());
582 }
583
584 #[test]
585 fn a_truncated_vector_is_refused_rather_than_read_past() {
586 let mut bytes = Vec::new();
587 vector(&(0..5000).map(|at| at % 3 == 0).collect::<Vec<bool>>()).write(&mut bytes);
588 assert!(BitVector::read(&bytes[..bytes.len() - 1], 5000).is_err());
589 assert!(BitVector::read(&bytes[..4], 5000).is_err());
590 }
591
592 #[test]
593 fn the_rank_index_costs_about_an_eighth_of_the_bitmap() {
594 let built = vector(&(0..1_000_000).map(|at| at % 5 == 0).collect::<Vec<bool>>());
597 let bitmap = 1_000_000 / 8;
598 assert!(built.bytes() > bitmap, "{} is not more than {bitmap}", built.bytes());
599 assert!(
600 built.bytes() < bitmap * 6 / 5,
601 "{} is more than a fifth over {bitmap}",
602 built.bytes()
603 );
604 }
605}