Skip to main content

vortex_mask/
intersect_by_rank.rs

1// SPDX-License-Identifier: Apache-2.0
2// SPDX-FileCopyrightText: Copyright the Vortex contributors
3
4use 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    /// Whether the implementation benefits from short-circuiting on `rank_bits == 0`
21    /// and `self_chunk == u64::MAX`. The portable path loops `popcount(mask)` times,
22    /// so an all-ones mask is genuinely expensive; BMI2 PDEP is constant-time and
23    /// the branches just add mispredict cost.
24    const PREFER_BRANCHES: bool;
25
26    fn deposit_bits(source: u64, mask: u64, mask_count: usize) -> u64;
27}
28
29trait SelectBit {
30    /// Position (0..63) of the `rank`-th set bit in `word`. Caller ensures
31    /// `rank < word.count_ones()`.
32    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        // SAFETY: callers only instantiate this implementation after checking BMI2 support.
119        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        // SAFETY: callers only instantiate this implementation after checking BMI2 support.
128        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    // PDEP places the rank-th bit of source into the rank-th set bit of mask, returning a single
145    // bit at the desired position.
146    let bit = x86_64::_pdep_u64(1u64 << rank, word);
147    bit.trailing_zeros() as usize
148}
149
150/// Reader that pulls variable-length (0..=64 bit) groups from a [`BitBuffer`] sequentially.
151///
152/// Maintains a 128-bit window over two consecutive chunks (`current`, `next`) and uses a
153/// funnel shift via `u128` to extract bits at any offset without branching. The shift
154/// pattern compiles to a single funnel-shift / SHRD-style sequence on x86_64.
155struct 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        // Funnel shift: extract `bit_count` bits at `bit_offset` from the (next:current)
188        // 128-bit window. For bit_offset in 0..=63 this is a single SHRD-style instruction
189        // on x86_64; the u128 cast keeps it well-defined when bit_offset == 0.
190        let combined = ((self.next as u128) << 64) | (self.current as u128);
191        // The truncation is intentional: we want the low 64 bits of the funnel-shifted
192        // window, which is exactly what `as u64` produces.
193        #[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    // SAFETY: callers allocate enough capacity for every output chunk.
258    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
324/// Walks `mask_indices` (global ranks into `self_buffer.set_bits`) and emits the corresponding
325/// positions in `self_buffer`. For each rank, advances `self_buffer`'s chunks via popcount
326/// skip-while, then locates the bit inside the current chunk with rank-select.
327///
328/// This dominates the chunk-scan paths when the mask is very sparse: cost is
329/// `O(mask.true_count() + self.len() / 64)` rather than `O(self.len() / 64)` per chunk.
330fn 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            // SAFETY: mask indices are ranks into self_indices, because
389            // mask.len() == self.true_count() == self_indices.len().
390            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
446/// Returns whether a mask is sparse.
447///
448/// [`BitBuffer`] traversal uses `u64` words, so fewer than one selected value per word is sparse.
449fn mask_is_sparse(values: &MaskValuesRef) -> bool {
450    values.true_count().saturating_mul(64) < values.len()
451}
452
453/// Returns whether a rank mask is sparse.
454///
455/// The mask-driven path becomes worthwhile at approximately 3% mask density. Each set bit costs a
456/// select and push, but this avoids a popcount and deposit for each chunk of `self`.
457fn rank_mask_is_sparse(values: &MaskValuesRef) -> bool {
458    values.true_count().saturating_mul(32) < values.len()
459}
460
461impl Mask {
462    /// Take the intersection of the `mask` with the set of true values in `self`.
463    ///
464    /// The hot path keeps bit-buffer-backed masks as bit buffers. It scans the set bits of `self`
465    /// by rank and deposits selected rank bits into their original positions.
466    ///
467    /// # Examples
468    ///
469    /// Keep the third and fifth set values from mask `m1`:
470    /// ```
471    /// use vortex_mask::Mask;
472    ///
473    /// let m1 = Mask::from_iter([true, false, false, true, true, true, false, true]);
474    /// let m2 = Mask::from_iter([false, false, true, false, true]);
475    /// assert_eq!(
476    ///     m1.intersect_by_rank(&m2),
477    ///     Mask::from_iter([false, false, false, false, true, false, false, true])
478    /// );
479    /// ```
480    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                // Four dispatch cases keyed by (self density, mask density):
489                //
490                //              | mask sparse | mask dense
491                // -------------+-------------+------------
492                // self sparse  | indices     | indices
493                // self dense   | mask-driven | bit-buffer
494                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        // Example from the documentation
654        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]); // 3 true values
665        let m2 = Mask::new_true(5); // 5 true values - doesn't match
666        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    // Larger sizes to push the bench-shaped buffer paths through the unit tests too.
708    #[case::dense_len_1024(1024, 31, 0.5, 0.5)]
709    // Very-sparse mask exercises the mask-driven dispatch path. Both densities live in
710    // the half-open interval where `mask_is_very_sparse` is true.
711    #[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}