Skip to main content

parquet/arrow/arrow_reader/selection/
algebra.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18//! Set algebra backing [`RowSelection::and_then`], [`RowSelection::intersection`]
19//! and [`RowSelection::union`]
20//!
21//! Each operation has two implementations, picked by the backing of its
22//! operands: a merge of the [`RowSelector`] runs, and a bitwise variant over
23//! [`BooleanBuffer`] masks.
24
25use super::{MaskRunIter, RowSelection, RowSelectionInner, RowSelector};
26use arrow_buffer::{BooleanBuffer, BooleanBufferBuilder};
27use std::cmp::Ordering;
28use std::iter::Peekable;
29
30/// Applies `second` to the rows selected by `first`, both selector-backed.
31pub(super) fn and_then_row_selections(
32    first: &[RowSelector],
33    second: &[RowSelector],
34) -> RowSelection {
35    let mut selectors = vec![];
36    let mut first = first.iter().copied().peekable();
37    let mut second = second.iter().copied().peekable();
38    and_then_iter(&mut selectors, &mut first, &mut second);
39    RowSelection::from_selectors(selectors)
40}
41
42/// Applies the mask `second` to the rows selected by the selector-backed `first`.
43///
44/// The mask is streamed as [`RowSelector`] runs, so it is never materialized.
45pub(super) fn and_then_selectors_with_mask(
46    first: &[RowSelector],
47    second: &BooleanBuffer,
48) -> RowSelection {
49    let mut selectors = vec![];
50    let mut first = first.iter().copied().peekable();
51    let mut second = MaskRunIter::new(second).peekable();
52    and_then_iter(&mut selectors, &mut first, &mut second);
53    RowSelection::from_selectors(selectors)
54}
55
56fn and_then_iter<I, J>(
57    selectors: &mut Vec<RowSelector>,
58    first: &mut Peekable<I>,
59    second: &mut Peekable<J>,
60) where
61    I: Iterator<Item = RowSelector>,
62    J: Iterator<Item = RowSelector>,
63{
64    let mut to_skip = 0;
65    while let Some(b) = second.peek_mut() {
66        let a = first
67            .peek_mut()
68            .expect("selection exceeds the number of selected rows");
69
70        if b.row_count == 0 {
71            second.next().unwrap();
72            continue;
73        }
74
75        if a.row_count == 0 {
76            first.next().unwrap();
77            continue;
78        }
79
80        if a.skip {
81            // Records were skipped when producing second
82            to_skip += a.row_count;
83            first.next().unwrap();
84            continue;
85        }
86
87        let skip = b.skip;
88        let to_process = a.row_count.min(b.row_count);
89
90        a.row_count -= to_process;
91        b.row_count -= to_process;
92
93        match skip {
94            true => to_skip += to_process,
95            false => {
96                if to_skip != 0 {
97                    selectors.push(RowSelector::skip(to_skip));
98                    to_skip = 0;
99                }
100                selectors.push(RowSelector::select(to_process))
101            }
102        }
103    }
104
105    for v in first {
106        if v.row_count != 0 {
107            assert!(
108                v.skip,
109                "selection contains less than the number of selected rows"
110            );
111            to_skip += v.row_count
112        }
113    }
114
115    if to_skip != 0 {
116        selectors.push(RowSelector::skip(to_skip));
117    }
118}
119
120/// Combine two lists of `RowSelection` return the intersection of them
121/// For example:
122/// self:      NNYYYYNNYYNYN
123/// other:     NYNNNNNNY
124///
125/// returned:  NNNNNNNNYYNYN
126pub(super) fn intersect_row_selections(
127    left: &[RowSelector],
128    right: &[RowSelector],
129) -> RowSelection {
130    let mut l_iter = left.iter().copied().peekable();
131    let mut r_iter = right.iter().copied().peekable();
132
133    let iter = std::iter::from_fn(move || {
134        loop {
135            let l = l_iter.peek_mut();
136            let r = r_iter.peek_mut();
137
138            match (l, r) {
139                (Some(a), _) if a.row_count == 0 => {
140                    l_iter.next().unwrap();
141                }
142                (_, Some(b)) if b.row_count == 0 => {
143                    r_iter.next().unwrap();
144                }
145                (Some(l), Some(r)) => {
146                    return match (l.skip, r.skip) {
147                        // Keep both ranges
148                        (false, false) => {
149                            if l.row_count < r.row_count {
150                                r.row_count -= l.row_count;
151                                l_iter.next()
152                            } else {
153                                l.row_count -= r.row_count;
154                                r_iter.next()
155                            }
156                        }
157                        // skip at least one
158                        _ => {
159                            if l.row_count < r.row_count {
160                                let skip = l.row_count;
161                                r.row_count -= l.row_count;
162                                l_iter.next();
163                                Some(RowSelector::skip(skip))
164                            } else {
165                                let skip = r.row_count;
166                                l.row_count -= skip;
167                                r_iter.next();
168                                Some(RowSelector::skip(skip))
169                            }
170                        }
171                    };
172                }
173                (Some(_), None) => return l_iter.next(),
174                (None, Some(_)) => return r_iter.next(),
175                (None, None) => return None,
176            }
177        }
178    });
179
180    iter.collect()
181}
182
183/// Combine two lists of `RowSelector` return the union of them
184/// For example:
185/// self:      NNYYYYNNYYNYN
186/// other:     NYNNNNNNY
187///
188/// returned:  NYYYYYNNYYNYN
189///
190/// This can be removed from here once RowSelection::union is in parquet::arrow
191pub(super) fn union_row_selections(left: &[RowSelector], right: &[RowSelector]) -> RowSelection {
192    let mut l_iter = left.iter().copied().peekable();
193    let mut r_iter = right.iter().copied().peekable();
194
195    let iter = std::iter::from_fn(move || {
196        loop {
197            let l = l_iter.peek_mut();
198            let r = r_iter.peek_mut();
199
200            match (l, r) {
201                (Some(a), _) if a.row_count == 0 => {
202                    l_iter.next().unwrap();
203                }
204                (_, Some(b)) if b.row_count == 0 => {
205                    r_iter.next().unwrap();
206                }
207                (Some(l), Some(r)) => {
208                    return match (l.skip, r.skip) {
209                        // Skip both ranges
210                        (true, true) => {
211                            if l.row_count < r.row_count {
212                                let skip = l.row_count;
213                                r.row_count -= l.row_count;
214                                l_iter.next();
215                                Some(RowSelector::skip(skip))
216                            } else {
217                                let skip = r.row_count;
218                                l.row_count -= skip;
219                                r_iter.next();
220                                Some(RowSelector::skip(skip))
221                            }
222                        }
223                        // Keep rows from left
224                        (false, true) => {
225                            if l.row_count < r.row_count {
226                                r.row_count -= l.row_count;
227                                l_iter.next()
228                            } else {
229                                let r_row_count = r.row_count;
230                                l.row_count -= r_row_count;
231                                r_iter.next();
232                                Some(RowSelector::select(r_row_count))
233                            }
234                        }
235                        // Keep rows from right
236                        (true, false) => {
237                            if l.row_count < r.row_count {
238                                let l_row_count = l.row_count;
239                                r.row_count -= l_row_count;
240                                l_iter.next();
241                                Some(RowSelector::select(l_row_count))
242                            } else {
243                                l.row_count -= r.row_count;
244                                r_iter.next()
245                            }
246                        }
247                        // Keep at least one
248                        _ => {
249                            if l.row_count < r.row_count {
250                                r.row_count -= l.row_count;
251                                l_iter.next()
252                            } else {
253                                l.row_count -= r.row_count;
254                                r_iter.next()
255                            }
256                        }
257                    };
258                }
259                (Some(_), None) => return l_iter.next(),
260                (None, Some(_)) => return r_iter.next(),
261                (None, None) => return None,
262            }
263        }
264    });
265
266    iter.collect()
267}
268
269/// Bitwise AND of two mask-backed selections. Longer side's tail passes through.
270pub(super) fn intersect_masks(l: &BooleanBuffer, r: &BooleanBuffer) -> BooleanBuffer {
271    if l.len() == r.len() {
272        return l & r;
273    }
274    let common = l.len().min(r.len());
275    let head = &l.slice(0, common) & &r.slice(0, common);
276    let (longer, longer_len) = if l.len() > r.len() {
277        (l, l.len())
278    } else {
279        (r, r.len())
280    };
281    let tail = longer.slice(common, longer_len - common);
282    let mut builder = BooleanBufferBuilder::new(longer_len);
283    builder.append_buffer(&head);
284    builder.append_buffer(&tail);
285    builder.finish()
286}
287
288/// Bitwise OR of two mask-backed selections. Longer side's tail passes through.
289pub(super) fn union_masks(l: &BooleanBuffer, r: &BooleanBuffer) -> BooleanBuffer {
290    if l.len() == r.len() {
291        return l | r;
292    }
293    let common = l.len().min(r.len());
294    let head = &l.slice(0, common) | &r.slice(0, common);
295    let (longer, longer_len) = if l.len() > r.len() {
296        (l, l.len())
297    } else {
298        (r, r.len())
299    };
300    let tail = longer.slice(common, longer_len - common);
301    let mut builder = BooleanBufferBuilder::new(longer_len);
302    builder.append_buffer(&head);
303    builder.append_buffer(&tail);
304    builder.finish()
305}
306
307/// Applies `other` to the selected rows of `mask`, preserving the original row domain.
308pub(super) fn and_then_mask(mask: &BooleanBuffer, other: &RowSelection) -> BooleanBuffer {
309    match &other.inner {
310        RowSelectionInner::Mask(other_mask) => and_then_masks(mask, other_mask.mask()),
311        RowSelectionInner::Selectors(selectors) => {
312            and_then_mask_from_selectors(mask, selectors.iter().copied())
313        }
314    }
315}
316
317fn and_then_mask_from_selectors<I>(mask: &BooleanBuffer, other: I) -> BooleanBuffer
318where
319    I: IntoIterator<Item = RowSelector>,
320{
321    let mut builder = BooleanBufferBuilder::new(mask.len());
322    let mut other_iter = other.into_iter();
323    let mut current = other_iter.next();
324    let mut cursor = 0usize;
325
326    // Iterate only over the set positions in `mask`; the gaps of unset bits
327    // are filled in bulk with `append_n` instead of bit-by-bit.
328    for set_idx in mask.set_indices() {
329        if set_idx > cursor {
330            builder.append_n(set_idx - cursor, false);
331        }
332        cursor = set_idx + 1;
333
334        while current.as_ref().is_some_and(|s| s.row_count == 0) {
335            current = other_iter.next();
336        }
337        let selector = current
338            .as_mut()
339            .expect("selection contains less than the number of selected rows");
340        let selected = !selector.skip;
341        selector.row_count -= 1;
342        builder.append(selected);
343    }
344    if cursor < mask.len() {
345        builder.append_n(mask.len() - cursor, false);
346    }
347
348    if current.is_some_and(|s| s.row_count != 0) || other_iter.any(|s| s.row_count != 0) {
349        panic!("selection exceeds the number of selected rows");
350    }
351
352    builder.finish()
353}
354
355fn and_then_masks(mask: &BooleanBuffer, other: &BooleanBuffer) -> BooleanBuffer {
356    let selected_count = mask.count_set_bits();
357    match other.len().cmp(&selected_count) {
358        Ordering::Less => panic!("selection contains less than the number of selected rows"),
359        Ordering::Greater => panic!("selection exceeds the number of selected rows"),
360        Ordering::Equal => {}
361    }
362
363    let other_true_count = other.count_set_bits();
364    if other_true_count == 0 {
365        return BooleanBuffer::new_unset(mask.len());
366    }
367    if other_true_count == selected_count {
368        return mask.clone();
369    }
370
371    let mut builder = BooleanBufferBuilder::new(mask.len());
372    let mut outer_set_indices = mask.set_indices();
373    let mut next_selected_ordinal = 0usize;
374    let mut cursor = 0usize;
375
376    for selected_ordinal in other.set_indices() {
377        let skip = selected_ordinal - next_selected_ordinal;
378        let set_idx = outer_set_indices
379            .nth(skip)
380            .expect("validated other length matches selected row count");
381        if set_idx > cursor {
382            builder.append_n(set_idx - cursor, false);
383        }
384        builder.append(true);
385        cursor = set_idx + 1;
386        next_selected_ordinal = selected_ordinal + 1;
387    }
388
389    if cursor < mask.len() {
390        builder.append_n(mask.len() - cursor, false);
391    }
392
393    builder.finish()
394}
395
396#[cfg(test)]
397mod tests {
398    use super::*;
399    use arrow_array::BooleanArray;
400    use rand::{Rng, rng};
401
402    #[test]
403    fn test_and() {
404        let mut a = RowSelection::from(vec![
405            RowSelector::skip(12),
406            RowSelector::select(23),
407            RowSelector::skip(3),
408            RowSelector::select(5),
409        ]);
410
411        let b = RowSelection::from(vec![
412            RowSelector::select(5),
413            RowSelector::skip(4),
414            RowSelector::select(15),
415            RowSelector::skip(4),
416        ]);
417
418        let mut expected = RowSelection::from(vec![
419            RowSelector::skip(12),
420            RowSelector::select(5),
421            RowSelector::skip(4),
422            RowSelector::select(14),
423            RowSelector::skip(3),
424            RowSelector::select(1),
425            RowSelector::skip(4),
426        ]);
427
428        assert_eq!(a.and_then(&b), expected);
429
430        a.split_off(7);
431        expected.split_off(7);
432        assert_eq!(a.and_then(&b), expected);
433
434        let a = RowSelection::from(vec![RowSelector::select(5), RowSelector::skip(3)]);
435
436        let b = RowSelection::from(vec![
437            RowSelector::select(2),
438            RowSelector::skip(1),
439            RowSelector::select(1),
440            RowSelector::skip(1),
441        ]);
442
443        assert_eq!(
444            a.and_then(&b).selectors(),
445            vec![
446                RowSelector::select(2),
447                RowSelector::skip(1),
448                RowSelector::select(1),
449                RowSelector::skip(4)
450            ]
451        );
452    }
453
454    #[test]
455    #[should_panic(expected = "selection exceeds the number of selected rows")]
456    fn test_and_longer() {
457        let a = RowSelection::from(vec![
458            RowSelector::select(3),
459            RowSelector::skip(33),
460            RowSelector::select(3),
461            RowSelector::skip(33),
462        ]);
463        let b = RowSelection::from(vec![RowSelector::select(36)]);
464        a.and_then(&b);
465    }
466
467    #[test]
468    #[should_panic(expected = "selection contains less than the number of selected rows")]
469    fn test_and_shorter() {
470        let a = RowSelection::from(vec![
471            RowSelector::select(3),
472            RowSelector::skip(33),
473            RowSelector::select(3),
474            RowSelector::skip(33),
475        ]);
476        let b = RowSelection::from(vec![RowSelector::select(3)]);
477        a.and_then(&b);
478    }
479
480    #[test]
481    fn test_intersect_row_selection_and_combine() {
482        // a size equal b size
483        let a = vec![
484            RowSelector::select(5),
485            RowSelector::skip(4),
486            RowSelector::select(1),
487        ];
488        let b = vec![
489            RowSelector::select(8),
490            RowSelector::skip(1),
491            RowSelector::select(1),
492        ];
493
494        let res = intersect_row_selections(&a, &b);
495        assert_eq!(
496            res.selectors(),
497            vec![
498                RowSelector::select(5),
499                RowSelector::skip(4),
500                RowSelector::select(1),
501            ],
502        );
503
504        // a size larger than b size
505        let a = vec![
506            RowSelector::select(3),
507            RowSelector::skip(33),
508            RowSelector::select(3),
509            RowSelector::skip(33),
510        ];
511        let b = vec![RowSelector::select(36), RowSelector::skip(36)];
512        let res = intersect_row_selections(&a, &b);
513        assert_eq!(
514            res.selectors(),
515            vec![RowSelector::select(3), RowSelector::skip(69)]
516        );
517
518        // a size less than b size
519        let a = vec![RowSelector::select(3), RowSelector::skip(7)];
520        let b = vec![
521            RowSelector::select(2),
522            RowSelector::skip(2),
523            RowSelector::select(2),
524            RowSelector::skip(2),
525            RowSelector::select(2),
526        ];
527        let res = intersect_row_selections(&a, &b);
528        assert_eq!(
529            res.selectors(),
530            vec![RowSelector::select(2), RowSelector::skip(8)]
531        );
532
533        let a = vec![RowSelector::select(3), RowSelector::skip(7)];
534        let b = vec![
535            RowSelector::select(2),
536            RowSelector::skip(2),
537            RowSelector::select(2),
538            RowSelector::skip(2),
539            RowSelector::select(2),
540        ];
541        let res = intersect_row_selections(&a, &b);
542        assert_eq!(
543            res.selectors(),
544            vec![RowSelector::select(2), RowSelector::skip(8)]
545        );
546    }
547
548    #[test]
549    fn test_and_fuzz() {
550        let mut rand = rng();
551        for _ in 0..100 {
552            let a_len = rand.random_range(10..100);
553            let a_bools: Vec<_> = (0..a_len).map(|_| rand.random_bool(0.2)).collect();
554            let a = RowSelection::from_filters(&[BooleanArray::from(a_bools.clone())]);
555
556            let b_len: usize = a_bools.iter().map(|x| *x as usize).sum();
557            let b_bools: Vec<_> = (0..b_len).map(|_| rand.random_bool(0.8)).collect();
558            let b = RowSelection::from_filters(&[BooleanArray::from(b_bools.clone())]);
559
560            let mut expected_bools = vec![false; a_len];
561
562            let mut iter_b = b_bools.iter();
563            for (idx, b) in a_bools.iter().enumerate() {
564                if *b && *iter_b.next().unwrap() {
565                    expected_bools[idx] = true;
566                }
567            }
568
569            let expected = RowSelection::from_filters(&[BooleanArray::from(expected_bools)]);
570
571            let total_rows: usize = expected.selectors().iter().map(|s| s.row_count).sum();
572            assert_eq!(a_len, total_rows);
573
574            assert_eq!(a.and_then(&b), expected);
575        }
576    }
577
578    #[test]
579    fn test_intersection() {
580        let selection = RowSelection::from(vec![RowSelector::select(1048576)]);
581        let result = selection.intersection(&selection);
582        assert_eq!(result, selection);
583
584        let a = RowSelection::from(vec![
585            RowSelector::skip(10),
586            RowSelector::select(10),
587            RowSelector::skip(10),
588            RowSelector::select(20),
589        ]);
590
591        let b = RowSelection::from(vec![
592            RowSelector::skip(20),
593            RowSelector::select(20),
594            RowSelector::skip(10),
595        ]);
596
597        let result = a.intersection(&b);
598        assert_eq!(
599            result.selectors(),
600            vec![
601                RowSelector::skip(30),
602                RowSelector::select(10),
603                RowSelector::skip(10)
604            ]
605        );
606    }
607
608    #[test]
609    fn test_union() {
610        let selection = RowSelection::from(vec![RowSelector::select(1048576)]);
611        let result = selection.union(&selection);
612        assert_eq!(result, selection);
613
614        // NYNYY
615        let a = RowSelection::from(vec![
616            RowSelector::skip(10),
617            RowSelector::select(10),
618            RowSelector::skip(10),
619            RowSelector::select(20),
620        ]);
621
622        // NNYYNYN
623        let b = RowSelection::from(vec![
624            RowSelector::skip(20),
625            RowSelector::select(20),
626            RowSelector::skip(10),
627            RowSelector::select(10),
628            RowSelector::skip(10),
629        ]);
630
631        let result = a.union(&b);
632
633        // NYYYYYN
634        assert_eq!(
635            result.iter().copied().collect::<Vec<_>>(),
636            vec![
637                RowSelector::skip(10),
638                RowSelector::select(50),
639                RowSelector::skip(10),
640            ]
641        );
642    }
643
644    #[test]
645    fn test_mask_and_then_preserves_backing() {
646        let outer_bits = vec![false, true, true, false, true, false, true];
647        let inner_bits = vec![true, false, true, false];
648        let outer_mask = RowSelection::from_boolean_buffer(BooleanBuffer::from(outer_bits.clone()));
649        let inner = RowSelection::from_filters(&[BooleanArray::from(inner_bits.clone())]);
650
651        let result = outer_mask.and_then(&inner);
652        assert!(result.as_mask().is_some());
653
654        let outer_selectors = RowSelection::from_filters(&[BooleanArray::from(outer_bits)]);
655        let expected = outer_selectors.and_then(&inner);
656        assert_eq!(result, expected);
657
658        let result_mask = result.as_mask().unwrap();
659        let actual_bits: Vec<_> = (0..result_mask.len())
660            .map(|i| result_mask.value(i))
661            .collect();
662        assert_eq!(
663            actual_bits,
664            vec![false, true, false, false, true, false, false]
665        );
666    }
667
668    #[test]
669    fn test_mask_and_then_mask_preserves_backing() {
670        let outer_bits = vec![false, true, true, false, true, false, true, true];
671        let inner_bits = vec![false, true, false, true, false];
672        let outer_mask = RowSelection::from_boolean_buffer(BooleanBuffer::from(outer_bits.clone()));
673        let inner_mask = RowSelection::from_boolean_buffer(BooleanBuffer::from(inner_bits));
674
675        let result = outer_mask.and_then(&inner_mask);
676        assert!(result.as_mask().is_some());
677
678        let outer_selectors = RowSelection::from_filters(&[BooleanArray::from(outer_bits)]);
679        let inner_selectors = RowSelection::from_filters(&[BooleanArray::from(vec![
680            false, true, false, true, false,
681        ])]);
682        assert_eq!(result, outer_selectors.and_then(&inner_selectors));
683
684        let result_mask = result.as_mask().unwrap();
685        let actual_bits: Vec<_> = (0..result_mask.len())
686            .map(|i| result_mask.value(i))
687            .collect();
688        assert_eq!(
689            actual_bits,
690            vec![false, false, true, false, false, false, true, false]
691        );
692    }
693
694    #[test]
695    fn test_selector_and_then_mask() {
696        let outer =
697            RowSelection::from_filters(&[BooleanArray::from(vec![false, true, true, false, true])]);
698        let inner = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![true, false, true]));
699
700        let result = outer.and_then(&inner);
701        assert!(result.as_mask().is_none());
702        assert_eq!(
703            result,
704            RowSelection::from_filters(&[BooleanArray::from(vec![
705                false, true, false, false, true,
706            ])])
707        );
708    }
709
710    #[test]
711    fn test_mask_and_then_none_selected_returns_all_unset() {
712        let outer = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![
713            false, true, true, false, true,
714        ]));
715        let inner =
716            RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![false, false, false]));
717
718        let result = outer.and_then(&inner);
719        let mask = result.as_mask().unwrap();
720        assert_eq!(mask.len(), 5);
721        assert_eq!(mask.count_set_bits(), 0);
722    }
723
724    #[test]
725    fn test_mask_intersection_uses_bitwise() {
726        let a_bits = vec![true, true, false, true, false, true];
727        let b_bits = vec![true, false, true, true, true, false];
728        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(a_bits.clone()));
729        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(b_bits.clone()));
730
731        let r = a.intersection(&b);
732        assert!(r.as_mask().is_some());
733
734        let expected: Vec<bool> = a_bits.iter().zip(&b_bits).map(|(x, y)| *x && *y).collect();
735        let expected_sel = RowSelection::from_filters(&[BooleanArray::from(expected)]);
736        assert_eq!(r, expected_sel);
737    }
738
739    #[test]
740    fn test_mask_union_uses_bitwise() {
741        let a_bits = vec![true, false, false, true, false, false];
742        let b_bits = vec![false, true, false, false, true, false];
743        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(a_bits.clone()));
744        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(b_bits.clone()));
745
746        let r = a.union(&b);
747        assert!(r.as_mask().is_some());
748
749        let expected: Vec<bool> = a_bits.iter().zip(&b_bits).map(|(x, y)| *x || *y).collect();
750        let expected_sel = RowSelection::from_filters(&[BooleanArray::from(expected)]);
751        assert_eq!(r, expected_sel);
752    }
753
754    #[test]
755    fn test_mixed_mask_selector_intersection_and_union() {
756        let mask_bits = vec![true, false, true, false, true, false];
757        let selector_bits = vec![false, true, true, false, false, true];
758        let mask = RowSelection::from_boolean_buffer(BooleanBuffer::from(mask_bits.clone()));
759        let selectors = RowSelection::from_filters(&[BooleanArray::from(selector_bits.clone())]);
760
761        let intersection_bits: Vec<_> = mask_bits
762            .iter()
763            .zip(&selector_bits)
764            .map(|(x, y)| *x && *y)
765            .collect();
766        let expected_intersection =
767            RowSelection::from_filters(&[BooleanArray::from(intersection_bits)]);
768        assert_eq!(mask.intersection(&selectors), expected_intersection);
769        assert_eq!(selectors.intersection(&mask), expected_intersection);
770
771        let union_bits: Vec<_> = mask_bits
772            .iter()
773            .zip(&selector_bits)
774            .map(|(x, y)| *x || *y)
775            .collect();
776        let expected_union = RowSelection::from_filters(&[BooleanArray::from(union_bits)]);
777        assert_eq!(mask.union(&selectors), expected_union);
778        assert_eq!(selectors.union(&mask), expected_union);
779    }
780
781    #[test]
782    fn test_mask_intersection_uneven_passes_tail_through() {
783        let a_bits = vec![true, true, true, true, true];
784        let b_bits = vec![true, false, true];
785        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(a_bits));
786        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(b_bits));
787
788        let r = a.intersection(&b);
789        let r_mask = r.as_mask().unwrap();
790        assert_eq!(r_mask.len(), 5);
791        let bits: Vec<bool> = (0..5).map(|i| r_mask.value(i)).collect();
792        assert_eq!(bits, vec![true, false, true, true, true]);
793
794        // Swapped operands: the right side is longer and its tail passes through.
795        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![true, false, true]));
796        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![
797            true, true, true, false, true,
798        ]));
799        let r = a.intersection(&b);
800        let r_mask = r.as_mask().unwrap();
801        assert_eq!(r_mask.len(), 5);
802        let bits: Vec<bool> = (0..5).map(|i| r_mask.value(i)).collect();
803        assert_eq!(bits, vec![true, false, true, false, true]);
804    }
805
806    #[test]
807    fn test_mask_union_uneven_passes_tail_through() {
808        let a_bits = vec![true, false, true];
809        let b_bits = vec![false, true, false, true, false];
810        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(a_bits));
811        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(b_bits));
812
813        let r = a.union(&b);
814        let r_mask = r.as_mask().unwrap();
815        assert_eq!(r_mask.len(), 5);
816        let bits: Vec<bool> = (0..5).map(|i| r_mask.value(i)).collect();
817        assert_eq!(bits, vec![true, true, true, true, false]);
818
819        let a = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![
820            false, true, false, false, true,
821        ]));
822        let b = RowSelection::from_boolean_buffer(BooleanBuffer::from(vec![true, false, false]));
823        let r = a.union(&b);
824        let r_mask = r.as_mask().unwrap();
825        let bits: Vec<bool> = (0..5).map(|i| r_mask.value(i)).collect();
826        assert_eq!(bits, vec![true, true, false, false, true]);
827    }
828}