1use std::iter::Chain;
5use std::iter::Once;
6use std::iter::once;
7use std::sync::Arc;
8
9use vortex_buffer::BitBuffer;
10use vortex_buffer::BitChunkIterator;
11use vortex_buffer::BufferMut;
12use vortex_buffer::CpuKernel;
13use vortex_error::VortexExpect;
14
15use crate::Mask;
16use crate::MaskValues;
17use crate::MaskValuesRef;
18
19trait DepositBits {
20 const PREFER_BRANCHES: bool;
25
26 fn deposit_bits(source: u64, mask: u64, mask_count: usize) -> u64;
27}
28
29trait SelectBit {
30 fn select_bit_position(word: u64, rank: usize) -> usize;
33}
34
35struct Portable;
36
37impl DepositBits for Portable {
38 const PREFER_BRANCHES: bool = true;
39
40 #[inline]
41 fn deposit_bits(source: u64, mask: u64, mask_count: usize) -> u64 {
42 if mask_count >= 16 && source.count_ones() as usize * 8 < mask_count {
43 return deposit_sparse_source(source, mask);
44 }
45
46 deposit_by_mask(source, mask)
47 }
48}
49
50impl SelectBit for Portable {
51 #[inline]
52 fn select_bit_position(word: u64, rank: usize) -> usize {
53 select_bit_position_portable(word, rank)
54 }
55}
56
57#[inline]
58fn deposit_by_mask(mut source: u64, mut mask: u64) -> u64 {
59 let mut result = 0u64;
60 while mask != 0 {
61 let bit = mask & mask.wrapping_neg();
62 if source & 1 != 0 {
63 result |= bit;
64 }
65 source >>= 1;
66 mask &= mask - 1;
67 }
68 result
69}
70
71#[inline]
72fn deposit_sparse_source(mut source: u64, mask: u64) -> u64 {
73 let mut result = 0u64;
74 while source != 0 {
75 result |= select_set_bit(mask, source.trailing_zeros() as usize);
76 source &= source - 1;
77 }
78 result
79}
80
81#[inline]
82fn select_set_bit(word: u64, rank: usize) -> u64 {
83 1u64 << select_bit_position_portable(word, rank)
84}
85
86#[inline]
87fn select_bit_position_portable(word: u64, mut rank: usize) -> usize {
88 debug_assert!(rank < word.count_ones() as usize);
89 let mut bit_offset = 0usize;
90 for byte in word.to_le_bytes() {
91 let count = byte.count_ones() as usize;
92 if rank < count {
93 let mut bits = byte;
94 for _ in 0..rank {
95 bits &= bits - 1;
96 }
97
98 return bit_offset + bits.trailing_zeros() as usize;
99 }
100
101 rank -= count;
102 bit_offset += 8;
103 }
104
105 debug_assert!(false, "rank out of bounds");
106 0
107}
108
109#[cfg(target_arch = "x86_64")]
110struct Bmi2;
111
112#[cfg(target_arch = "x86_64")]
113impl DepositBits for Bmi2 {
114 const PREFER_BRANCHES: bool = false;
115
116 #[inline]
117 fn deposit_bits(source: u64, mask: u64, _mask_count: usize) -> u64 {
118 unsafe { pdep_bmi2(source, mask) }
120 }
121}
122
123#[cfg(target_arch = "x86_64")]
124impl SelectBit for Bmi2 {
125 #[inline]
126 fn select_bit_position(word: u64, rank: usize) -> usize {
127 unsafe { select_bit_position_bmi2(word, rank) }
129 }
130}
131
132#[cfg(target_arch = "x86_64")]
133#[target_feature(enable = "bmi2")]
134unsafe fn pdep_bmi2(source: u64, mask: u64) -> u64 {
135 use std::arch::x86_64;
136 x86_64::_pdep_u64(source, mask)
137}
138
139#[cfg(target_arch = "x86_64")]
140#[target_feature(enable = "bmi2")]
141unsafe fn select_bit_position_bmi2(word: u64, rank: usize) -> usize {
142 use std::arch::x86_64;
143 debug_assert!(rank < word.count_ones() as usize);
144 let bit = x86_64::_pdep_u64(1u64 << rank, word);
147 bit.trailing_zeros() as usize
148}
149
150struct RankBitReader<'a> {
156 chunk_iter: Chain<BitChunkIterator<'a>, Once<u64>>,
157 current: u64,
158 next: u64,
159 bit_offset: usize,
160}
161
162impl<'a> RankBitReader<'a> {
163 fn new(buffer: &'a BitBuffer) -> Self {
164 let chunks = buffer.chunks();
165 let mut chunk_iter = chunks.iter().chain(once(chunks.remainder_bits()));
166
167 let current = chunk_iter.next().unwrap_or(0);
168 let next = chunk_iter.next().unwrap_or(0);
169
170 Self {
171 chunk_iter,
172 current,
173 next,
174 bit_offset: 0,
175 }
176 }
177
178 #[inline]
179 fn fetch_next(&mut self) -> u64 {
180 self.chunk_iter.next().unwrap_or(0)
181 }
182
183 #[inline]
184 fn read(&mut self, bit_count: usize) -> u64 {
185 debug_assert!(bit_count <= 64);
186
187 let combined = ((self.next as u128) << 64) | (self.current as u128);
191 #[expect(clippy::cast_possible_truncation)]
194 let bits = (combined >> self.bit_offset) as u64 & low_bits(bit_count);
195
196 let new_offset = self.bit_offset + bit_count;
197 if new_offset >= 64 {
198 self.current = self.next;
199 self.next = self.fetch_next();
200 self.bit_offset = new_offset - 64;
201 } else {
202 self.bit_offset = new_offset;
203 }
204
205 bits
206 }
207}
208
209#[inline]
210fn low_bits(bit_count: usize) -> u64 {
211 debug_assert!(bit_count <= 64);
212 if bit_count == 64 {
213 u64::MAX
214 } else {
215 (1u64 << bit_count) - 1
216 }
217}
218
219#[inline]
220fn mask_from_buffer(buffer: BitBuffer, true_count: usize) -> Mask {
221 let len = buffer.len();
222 if true_count == 0 {
223 return Mask::new_false(len);
224 }
225 if true_count == len {
226 return Mask::new_true(len);
227 }
228
229 Mask::Values(Arc::new(MaskValues {
230 buffer,
231 indices: Default::default(),
232 slices: Default::default(),
233 true_count,
234 density: true_count as f64 / len as f64,
235 }))
236}
237
238#[inline]
239fn push_result_chunk<D: DepositBits>(
240 result: &mut BufferMut<u64>,
241 self_chunk: u64,
242 self_count: usize,
243 rank_bits: u64,
244) {
245 let chunk = if D::PREFER_BRANCHES {
246 if rank_bits == 0 {
247 0
248 } else if self_chunk == u64::MAX {
249 rank_bits
250 } else {
251 D::deposit_bits(rank_bits, self_chunk, self_count)
252 }
253 } else {
254 D::deposit_bits(rank_bits, self_chunk, self_count)
255 };
256
257 unsafe { result.push_unchecked(chunk) };
259}
260
261fn intersect_bit_buffers<D: DepositBits>(
262 self_buffer: &BitBuffer,
263 mask_buffer: &BitBuffer,
264 true_count: usize,
265) -> Mask {
266 let len = self_buffer.len();
267 let mut result = BufferMut::with_capacity(len.div_ceil(64));
268 let mut reader = RankBitReader::new(mask_buffer);
269 let self_chunks = self_buffer.chunks();
270
271 for self_chunk in self_chunks.iter() {
272 let self_count = self_chunk.count_ones() as usize;
273 let rank_bits = reader.read(self_count);
274 push_result_chunk::<D>(&mut result, self_chunk, self_count, rank_bits);
275 }
276
277 if self_chunks.remainder_len() != 0 {
278 let self_chunk = self_chunks.remainder_bits();
279 let self_count = self_chunk.count_ones() as usize;
280 let rank_bits = reader.read(self_count);
281 push_result_chunk::<D>(&mut result, self_chunk, self_count, rank_bits);
282 }
283
284 mask_from_buffer(
285 BitBuffer::new(result.freeze().into_byte_buffer(), len),
286 true_count,
287 )
288}
289
290fn intersect_bit_buffer_by_rank_indices<D: DepositBits>(
291 self_buffer: &BitBuffer,
292 mask_indices: &[usize],
293) -> Mask {
294 let len = self_buffer.len();
295 let mut result = BufferMut::with_capacity(len.div_ceil(64));
296 let self_chunks = self_buffer.chunks();
297 let mut rank_base = 0usize;
298 let mut rank_idx = 0usize;
299
300 for self_chunk in self_chunks.iter() {
301 let self_count = self_chunk.count_ones() as usize;
302 let next_rank_base = rank_base + self_count;
303 let rank_bits = rank_bits_for_chunk(mask_indices, &mut rank_idx, rank_base, next_rank_base);
304 push_result_chunk::<D>(&mut result, self_chunk, self_count, rank_bits);
305 rank_base = next_rank_base;
306 }
307
308 if self_chunks.remainder_len() != 0 {
309 let self_chunk = self_chunks.remainder_bits();
310 let self_count = self_chunk.count_ones() as usize;
311 let next_rank_base = rank_base + self_count;
312 let rank_bits = rank_bits_for_chunk(mask_indices, &mut rank_idx, rank_base, next_rank_base);
313 push_result_chunk::<D>(&mut result, self_chunk, self_count, rank_bits);
314 }
315
316 debug_assert_eq!(rank_idx, mask_indices.len());
317
318 mask_from_buffer(
319 BitBuffer::new(result.freeze().into_byte_buffer(), len),
320 mask_indices.len(),
321 )
322}
323
324fn intersect_mask_driven<S, I>(self_buffer: &BitBuffer, mask_indices: I, true_count: usize) -> Mask
331where
332 S: SelectBit,
333 I: Iterator<Item = usize>,
334{
335 let len = self_buffer.len();
336 if true_count == 0 {
337 return Mask::new_false(len);
338 }
339
340 let mut chunk_iter = self_buffer.chunks().iter_padded();
341
342 let mut current_chunk = chunk_iter.next().unwrap_or(0);
343 let mut current_count = current_chunk.count_ones() as usize;
344 let mut current_chunk_idx = 0usize;
345 let mut rank_before = 0usize;
346
347 let mut output = Vec::with_capacity(true_count);
348
349 for global_rank in mask_indices {
350 while rank_before + current_count <= global_rank {
351 rank_before += current_count;
352 current_chunk_idx += 1;
353 current_chunk = chunk_iter.next().vortex_expect("mask index out of bounds");
354 current_count = current_chunk.count_ones() as usize;
355 }
356
357 let local_rank = global_rank - rank_before;
358 let bit_pos = S::select_bit_position(current_chunk, local_rank);
359 output.push(current_chunk_idx * 64 + bit_pos);
360 }
361
362 debug_assert_eq!(output.len(), true_count);
363 Mask::from_indices(len, output)
364}
365
366#[inline]
367fn rank_bits_for_chunk(
368 mask_indices: &[usize],
369 rank_idx: &mut usize,
370 rank_base: usize,
371 next_rank_base: usize,
372) -> u64 {
373 let mut rank_bits = 0u64;
374 while let Some(&rank) = mask_indices.get(*rank_idx) {
375 if rank >= next_rank_base {
376 break;
377 }
378 rank_bits |= 1u64 << (rank - rank_base);
379 *rank_idx += 1;
380 }
381 rank_bits
382}
383
384fn intersect_by_rank_indices(len: usize, self_indices: &[usize], mask_indices: &[usize]) -> Mask {
385 Mask::from_indices(
386 len,
387 mask_indices.iter().map(|idx| {
388 unsafe { *self_indices.get_unchecked(*idx) }
391 }),
392 )
393}
394
395#[inline]
396fn intersect_bit_buffers_dispatch(
397 self_buffer: &BitBuffer,
398 mask_buffer: &BitBuffer,
399 true_count: usize,
400) -> Mask {
401 type IntersectBuffers = fn(&BitBuffer, &BitBuffer, usize) -> Mask;
402 static KERNEL: CpuKernel<IntersectBuffers> = CpuKernel::new(|| {
403 #[cfg(target_arch = "x86_64")]
404 {
405 if std::arch::is_x86_feature_detected!("bmi2") {
406 return intersect_bit_buffers::<Bmi2>;
407 }
408 }
409 intersect_bit_buffers::<Portable>
410 });
411 KERNEL.get()(self_buffer, mask_buffer, true_count)
412}
413
414#[inline]
415fn intersect_rank_indices_dispatch(self_buffer: &BitBuffer, mask_indices: &[usize]) -> Mask {
416 type IntersectRankIndices = fn(&BitBuffer, &[usize]) -> Mask;
417 static KERNEL: CpuKernel<IntersectRankIndices> = CpuKernel::new(|| {
418 #[cfg(target_arch = "x86_64")]
419 {
420 if std::arch::is_x86_feature_detected!("bmi2") {
421 return intersect_bit_buffer_by_rank_indices::<Bmi2>;
422 }
423 }
424 intersect_bit_buffer_by_rank_indices::<Portable>
425 });
426 KERNEL.get()(self_buffer, mask_indices)
427}
428
429#[inline]
430fn intersect_mask_driven_dispatch<I>(
431 self_buffer: &BitBuffer,
432 mask_indices: I,
433 true_count: usize,
434) -> Mask
435where
436 I: Iterator<Item = usize>,
437{
438 #[cfg(target_arch = "x86_64")]
439 if std::arch::is_x86_feature_detected!("bmi2") {
440 return intersect_mask_driven::<Bmi2, _>(self_buffer, mask_indices, true_count);
441 }
442
443 intersect_mask_driven::<Portable, _>(self_buffer, mask_indices, true_count)
444}
445
446fn mask_is_sparse(values: &MaskValuesRef) -> bool {
450 values.true_count().saturating_mul(64) < values.len()
451}
452
453fn rank_mask_is_sparse(values: &MaskValuesRef) -> bool {
458 values.true_count().saturating_mul(32) < values.len()
459}
460
461impl Mask {
462 pub fn intersect_by_rank(&self, mask: &Mask) -> Mask {
481 assert_eq!(self.true_count(), mask.len());
482
483 match (self, mask) {
484 (Self::AllTrue(_), _) => mask.clone(),
485 (_, Self::AllTrue(_)) => self.clone(),
486 (Self::AllFalse(_), _) | (_, Self::AllFalse(_)) => Self::new_false(self.len()),
487 (Self::Values(self_values), Self::Values(mask_values)) => {
488 if let Some(mask_indices) = mask_values.indices.get() {
495 if let Some(self_indices) = self_values.indices.get()
496 && mask_indices.len() < self.len().div_ceil(64)
497 {
498 return intersect_by_rank_indices(self.len(), self_indices, mask_indices);
499 }
500
501 let self_is_very_sparse = mask_is_sparse(self_values);
502 let mask_is_very_sparse = rank_mask_is_sparse(mask_values);
503
504 if self_is_very_sparse {
505 return intersect_by_rank_indices(
506 self.len(),
507 self_values.indices(),
508 mask_indices,
509 );
510 }
511
512 if mask_is_very_sparse {
513 return intersect_mask_driven_dispatch(
514 self_values.bit_buffer(),
515 mask_indices.iter().copied(),
516 mask_values.true_count(),
517 );
518 }
519
520 if mask_indices.len().saturating_mul(4) > mask.len() {
521 return intersect_bit_buffers_dispatch(
522 self_values.bit_buffer(),
523 mask_values.bit_buffer(),
524 mask_values.true_count(),
525 );
526 }
527
528 return intersect_rank_indices_dispatch(self_values.bit_buffer(), mask_indices);
529 }
530
531 let self_is_very_sparse = mask_is_sparse(self_values);
532 let mask_is_very_sparse = rank_mask_is_sparse(mask_values);
533
534 if self_is_very_sparse {
535 return intersect_by_rank_indices(
536 self.len(),
537 self_values.indices(),
538 mask_values.indices(),
539 );
540 }
541
542 if mask_is_very_sparse {
543 return intersect_mask_driven_dispatch(
544 self_values.bit_buffer(),
545 mask_values.bit_buffer().set_indices(),
546 mask_values.true_count(),
547 );
548 }
549
550 intersect_bit_buffers_dispatch(
551 self_values.bit_buffer(),
552 mask_values.bit_buffer(),
553 mask_values.true_count(),
554 )
555 }
556 }
557 }
558}
559
560#[cfg(test)]
561mod tests {
562 use rstest::rstest;
563 use vortex_buffer::BitBuffer;
564
565 use crate::Mask;
566
567 #[test]
568 fn mask_bitand_all_as_bit_and() {
569 let this = Mask::from_buffer(BitBuffer::from_iter(vec![true, true, true, true, true]));
570 let mask = Mask::from_buffer(BitBuffer::from_iter(vec![false, true, false, true, true]));
571 assert_eq!(
572 this.intersect_by_rank(&mask),
573 Mask::from_indices(5, vec![1, 3, 4])
574 );
575 }
576
577 #[test]
578 fn mask_bitand_all_true() {
579 let this = Mask::from_buffer(BitBuffer::from_iter(vec![false, false, true, true, true]));
580 let mask = Mask::from_buffer(BitBuffer::from_iter(vec![true, true, true]));
581 assert_eq!(
582 this.intersect_by_rank(&mask),
583 Mask::from_indices(5, vec![2, 3, 4])
584 );
585 }
586
587 #[test]
588 fn mask_bitand_true() {
589 let this = Mask::from_buffer(BitBuffer::from_iter(vec![true, false, false, true, true]));
590 let mask = Mask::from_buffer(BitBuffer::from_iter(vec![true, false, true]));
591 assert_eq!(
592 this.intersect_by_rank(&mask),
593 Mask::from_indices(5, vec![0, 4])
594 );
595 }
596
597 #[test]
598 fn mask_bitand_false() {
599 let this = Mask::from_buffer(BitBuffer::from_iter(vec![true, false, false, true, true]));
600 let mask = Mask::from_buffer(BitBuffer::from_iter(vec![false, false, false]));
601 assert_eq!(this.intersect_by_rank(&mask), Mask::from_indices(5, vec![]));
602 }
603
604 #[test]
605 fn mask_intersect_by_rank_all_false() {
606 let this = Mask::AllFalse(10);
607 let mask = Mask::AllFalse(0);
608 assert_eq!(this.intersect_by_rank(&mask), Mask::AllFalse(10));
609 }
610
611 #[rstest]
612 #[case::all_true_with_all_true(
613 Mask::new_true(5),
614 Mask::new_true(5),
615 vec![0, 1, 2, 3, 4]
616 )]
617 #[case::all_true_with_all_false(
618 Mask::new_true(5),
619 Mask::new_false(5),
620 vec![]
621 )]
622 #[case::all_false_with_any(
623 Mask::new_false(10),
624 Mask::new_true(0),
625 vec![]
626 )]
627 #[case::indices_with_all_true(
628 Mask::from_indices(10, vec![2, 5, 7, 9]),
629 Mask::new_true(4),
630 vec![2, 5, 7, 9]
631 )]
632 #[case::indices_with_all_false(
633 Mask::from_indices(10, vec![2, 5, 7, 9]),
634 Mask::new_false(4),
635 vec![]
636 )]
637 fn test_intersect_by_rank_special_cases(
638 #[case] base_mask: Mask,
639 #[case] rank_mask: Mask,
640 #[case] expected_indices: Vec<usize>,
641 ) {
642 let result = base_mask.intersect_by_rank(&rank_mask);
643
644 match result.indices() {
645 crate::AllOr::All => assert_eq!(expected_indices.len(), result.len()),
646 crate::AllOr::None => assert!(expected_indices.is_empty()),
647 crate::AllOr::Some(indices) => assert_eq!(indices, &expected_indices[..]),
648 }
649 }
650
651 #[test]
652 fn test_intersect_by_rank_example() {
653 let m1 = Mask::from_iter([true, false, false, true, true, true, false, true]);
655 let m2 = Mask::from_iter([false, false, true, false, true]);
656 let result = m1.intersect_by_rank(&m2);
657 let expected = Mask::from_iter([false, false, false, false, true, false, false, true]);
658 assert_eq!(result, expected);
659 }
660
661 #[test]
662 #[should_panic]
663 fn test_intersect_by_rank_wrong_length() {
664 let m1 = Mask::from_indices(10, vec![2, 5, 7]); let m2 = Mask::new_true(5); m1.intersect_by_rank(&m2);
667 }
668
669 #[rstest]
670 #[case::single_element(
671 vec![3],
672 vec![true],
673 vec![3]
674 )]
675 #[case::single_element_masked(
676 vec![3],
677 vec![false],
678 vec![]
679 )]
680 #[case::alternating(
681 vec![0, 2, 4, 6, 8],
682 vec![true, false, true, false, true],
683 vec![0, 4, 8]
684 )]
685 #[case::consecutive(
686 vec![5, 6, 7, 8, 9],
687 vec![false, true, true, true, false],
688 vec![6, 7, 8]
689 )]
690 fn test_intersect_by_rank_patterns(
691 #[case] base_indices: Vec<usize>,
692 #[case] rank_pattern: Vec<bool>,
693 #[case] expected_indices: Vec<usize>,
694 ) {
695 let base = Mask::from_indices(10, base_indices);
696 let rank = Mask::from_iter(rank_pattern);
697 let result = base.intersect_by_rank(&rank);
698
699 match result.indices() {
700 crate::AllOr::Some(indices) => assert_eq!(indices, &expected_indices[..]),
701 crate::AllOr::None => assert!(expected_indices.is_empty()),
702 _ => panic!("Unexpected result"),
703 }
704 }
705
706 #[rstest]
707 #[case::dense_len_1024(1024, 31, 0.5, 0.5)]
709 #[case::sparse_mask_1pct(1024, 17, 0.5, 0.01)]
712 #[case::sparse_mask_2pct(2048, 0, 0.5, 0.02)]
713 #[case::very_sparse_mask_with_offsets(513, 5, 0.5, 0.005)]
714 fn test_intersect_by_rank_density_matrix(
715 #[case] base_len: usize,
716 #[case] base_offset: usize,
717 #[case] base_density: f64,
718 #[case] rank_density: f64,
719 ) {
720 #[expect(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
721 let base_threshold = (base_density * 1024.0) as usize;
722 #[expect(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
723 let rank_threshold = (rank_density * 1024.0) as usize;
724
725 let base_source: Vec<bool> = (0..base_len + base_offset + 16)
726 .map(|i| (i * 7 + 13) % 1024 < base_threshold)
727 .collect();
728 let base_bits = base_source[base_offset..base_offset + base_len].to_vec();
729 let base = Mask::from_buffer(
730 BitBuffer::from(base_source).slice(base_offset..base_offset + base_len),
731 );
732
733 let rank_len = base.true_count();
734 let rank_bits: Vec<bool> = (0..rank_len)
735 .map(|i| (i * 11 + 7) % 1024 < rank_threshold)
736 .collect();
737 let rank_from_buffer = Mask::from_buffer(BitBuffer::from(rank_bits.clone()));
738 let rank_indices_vec = rank_bits
739 .iter()
740 .enumerate()
741 .filter_map(|(idx, &v)| v.then_some(idx))
742 .collect::<Vec<_>>();
743 let rank_from_indices = Mask::from_indices(rank_len, rank_indices_vec);
744
745 let expected = expected_intersect_by_rank(&base_bits, &rank_bits);
746
747 assert_eq!(
748 base.intersect_by_rank(&rank_from_buffer),
749 expected,
750 "uncached rank"
751 );
752 assert_eq!(
753 base.intersect_by_rank(&rank_from_indices),
754 expected,
755 "cached rank"
756 );
757 }
758
759 #[rstest]
760 #[case::short(37, 0, 0)]
761 #[case::base_offset(257, 5, 0)]
762 #[case::rank_offset(257, 0, 3)]
763 #[case::both_offsets(513, 6, 5)]
764 fn test_intersect_by_rank_bitbuffer_paths_with_offsets(
765 #[case] base_len: usize,
766 #[case] base_offset: usize,
767 #[case] rank_offset: usize,
768 ) {
769 let base_source: Vec<bool> = (0..base_len + base_offset + 16)
770 .map(|i| (i % 3 == 0) ^ (i % 11 == 0) ^ (i % 17 == 0))
771 .collect();
772 let base_bits = base_source[base_offset..base_offset + base_len].to_vec();
773 let base = Mask::from_buffer(
774 BitBuffer::from(base_source).slice(base_offset..base_offset + base_len),
775 );
776
777 let rank_len = base.true_count();
778 let rank_bits: Vec<bool> = (0..rank_len)
779 .map(|i| (i % 5 == 0) || (i % 13 == 3))
780 .collect();
781 let mut rank_source = vec![false; rank_offset];
782 rank_source.extend(rank_bits.iter().copied());
783 rank_source.extend([true, false, true, false, true, false, true, false]);
784
785 let rank_from_buffer = Mask::from_buffer(
786 BitBuffer::from(rank_source).slice(rank_offset..rank_offset + rank_len),
787 );
788 let rank_indices = rank_bits
789 .iter()
790 .enumerate()
791 .filter_map(|(idx, &value)| value.then_some(idx))
792 .collect::<Vec<_>>();
793 let rank_from_indices = Mask::from_indices(rank_len, rank_indices);
794
795 let expected = expected_intersect_by_rank(&base_bits, &rank_bits);
796
797 assert_eq!(base.intersect_by_rank(&rank_from_buffer), expected);
798 assert_eq!(base.intersect_by_rank(&rank_from_indices), expected);
799 }
800
801 fn expected_intersect_by_rank(base_bits: &[bool], rank_bits: &[bool]) -> Mask {
802 let mut rank = 0usize;
803 Mask::from_iter(base_bits.iter().map(|&is_set| {
804 if is_set {
805 let keep = rank_bits[rank];
806 rank += 1;
807 keep
808 } else {
809 false
810 }
811 }))
812 }
813}