1use std::ops::AddAssign;
21use std::sync::Arc;
22
23use arrow_array::builder::BooleanBufferBuilder;
24use arrow_array::cast::AsArray;
25use arrow_array::types::{
26 ArrowDictionaryKeyType, ArrowPrimitiveType, ByteArrayType, ByteViewType, RunEndIndexType,
27};
28use arrow_array::*;
29use arrow_buffer::{
30 ArrowNativeType, BooleanBuffer, NullBuffer, OffsetBuffer, RunEndBuffer, ScalarBuffer, bit_util,
31};
32use arrow_buffer::{Buffer, MutableBuffer};
33use arrow_data::bit_iterator::{BitIndexIterator, BitSliceIterator};
34use arrow_data::transform::MutableArrayData;
35use arrow_schema::*;
36
37const FILTER_SLICES_SELECTIVITY_THRESHOLD: f64 = 0.8;
44
45#[derive(Debug)]
57pub struct SlicesIterator<'a>(BitSliceIterator<'a>);
58
59impl<'a> SlicesIterator<'a> {
60 pub fn new(filter: &'a BooleanArray) -> Self {
62 filter.values().into()
63 }
64}
65
66impl<'a> From<&'a BooleanBuffer> for SlicesIterator<'a> {
67 fn from(filter: &'a BooleanBuffer) -> Self {
68 Self(filter.set_slices())
69 }
70}
71
72impl Iterator for SlicesIterator<'_> {
73 type Item = (usize, usize);
74
75 fn next(&mut self) -> Option<Self::Item> {
76 self.0.next()
77 }
78}
79
80pub(crate) struct IndexIterator<'a> {
85 remaining: usize,
86 iter: BitIndexIterator<'a>,
87}
88
89impl<'a> IndexIterator<'a> {
90 pub(crate) fn new(filter: &'a BooleanArray, remaining: usize) -> Self {
91 assert_eq!(filter.null_count(), 0);
92 let iter = filter.values().set_indices();
93 Self { remaining, iter }
94 }
95
96 pub fn collect(mut self) -> Vec<usize> {
100 let len = self.remaining;
101 let mut result = Vec::with_capacity(len);
102 let ptr: *mut usize = result.as_mut_ptr();
103 for i in 0..len {
104 let next = self.iter.next();
107 debug_assert!(next.is_some(), "IndexIterator exhausted early");
108 unsafe {
109 *ptr.add(i) = next.unwrap_unchecked();
110 }
111 }
112 unsafe {
114 result.set_len(len);
115 }
116 result
117 }
118}
119
120impl Iterator for IndexIterator<'_> {
121 type Item = usize;
122
123 fn next(&mut self) -> Option<Self::Item> {
124 if self.remaining != 0 {
125 let next = self.iter.next().expect("IndexIterator exhausted early");
128 self.remaining -= 1;
129 return Some(next);
131 }
132 None
133 }
134
135 fn size_hint(&self) -> (usize, Option<usize>) {
136 (self.remaining, Some(self.remaining))
137 }
138}
139
140pub fn prep_null_mask_filter(filter: &BooleanArray) -> BooleanArray {
168 let nulls = filter.nulls().unwrap();
169 let mask = filter.values() & nulls.inner();
170 BooleanArray::new(mask, None)
171}
172
173pub fn filter(values: &dyn Array, predicate: &BooleanArray) -> Result<ArrayRef, ArrowError> {
202 let mut filter_builder = FilterBuilder::new(predicate);
203
204 if FilterBuilder::is_optimize_beneficial(values.data_type()) {
205 filter_builder = filter_builder.optimize();
208 }
209
210 let predicate = filter_builder.build();
211
212 filter_array(values, &predicate)
213}
214
215pub fn filter_record_batch(
226 record_batch: &RecordBatch,
227 predicate: &BooleanArray,
228) -> Result<RecordBatch, ArrowError> {
229 let mut filter_builder = FilterBuilder::new(predicate);
230 let num_cols = record_batch.num_columns();
231 if num_cols > 1
232 || (num_cols > 0
233 && FilterBuilder::is_optimize_beneficial(
234 record_batch.schema_ref().field(0).data_type(),
235 ))
236 {
237 filter_builder = filter_builder.optimize();
240 }
241 let filter = filter_builder.build();
242
243 filter.filter_record_batch(record_batch)
244}
245
246#[derive(Debug)]
248pub struct FilterBuilder {
249 filter: BooleanArray,
250 count: usize,
251 strategy: IterationStrategy,
252}
253
254impl FilterBuilder {
255 pub fn new(filter: &BooleanArray) -> Self {
257 Self::new_with_count(filter, filter.true_count())
258 }
259
260 pub(crate) fn new_with_count(filter: &BooleanArray, count: usize) -> Self {
261 let filter = match filter.null_count() {
262 0 => filter.clone(),
263 _ => prep_null_mask_filter(filter),
264 };
265
266 let strategy = IterationStrategy::default_strategy(filter.len(), count);
267
268 Self {
269 filter,
270 count,
271 strategy,
272 }
273 }
274
275 pub fn optimize(mut self) -> Self {
286 match self.strategy {
287 IterationStrategy::SlicesIterator => {
288 let slices = SlicesIterator::new(&self.filter).collect();
289 self.strategy = IterationStrategy::Slices(slices)
290 }
291 IterationStrategy::IndexIterator => {
292 let indices = IndexIterator::new(&self.filter, self.count).collect();
293 self.strategy = IterationStrategy::Indices(indices)
294 }
295 _ => {}
296 }
297 self
298 }
299
300 pub fn is_optimize_beneficial(data_type: &DataType) -> bool {
305 match data_type {
306 DataType::Struct(fields) => {
307 fields.len() > 1
308 || fields.len() == 1
309 && FilterBuilder::is_optimize_beneficial(fields[0].data_type())
310 }
311 DataType::Union(fields, UnionMode::Sparse) => !fields.is_empty(),
312 _ => false,
313 }
314 }
315
316 pub fn build(self) -> FilterPredicate {
318 FilterPredicate {
319 filter: self.filter,
320 count: self.count,
321 strategy: self.strategy,
322 }
323 }
324}
325
326#[derive(Debug)]
328enum IterationStrategy {
329 SlicesIterator,
331 IndexIterator,
333 Indices(Vec<usize>),
335 Slices(Vec<(usize, usize)>),
337 All,
339 None,
341}
342
343impl IterationStrategy {
344 fn default_strategy(filter_length: usize, filter_count: usize) -> Self {
347 if filter_length == 0 || filter_count == 0 {
348 return IterationStrategy::None;
349 }
350
351 if filter_count == filter_length {
352 return IterationStrategy::All;
353 }
354
355 let selectivity_frac = filter_count as f64 / filter_length as f64;
360 if selectivity_frac > FILTER_SLICES_SELECTIVITY_THRESHOLD {
361 return IterationStrategy::SlicesIterator;
362 }
363 IterationStrategy::IndexIterator
364 }
365}
366
367pub(crate) enum FilterSelection<'a> {
373 None,
375 All { len: usize },
377 Slices(FilterSlices<'a>),
379 Indices(FilterIndices<'a>),
381}
382
383pub(crate) type FilterSlices<'a> =
384 FilterIterator<std::iter::Copied<std::slice::Iter<'a, (usize, usize)>>, SlicesIterator<'a>>;
385
386pub(crate) type FilterIndices<'a> =
387 FilterIterator<std::iter::Copied<std::slice::Iter<'a, usize>>, IndexIterator<'a>>;
388
389pub(crate) enum FilterIterator<M, I> {
397 Materialized(M),
398 Lazy(I),
399}
400
401impl<M, I> FilterIterator<M, I>
402where
403 M: Iterator,
404 I: Iterator<Item = M::Item>,
405{
406 pub(crate) fn for_each<F>(self, f: F)
408 where
409 F: FnMut(M::Item),
410 {
411 match self {
412 Self::Materialized(iter) => iter.for_each(f),
413 Self::Lazy(iter) => iter.for_each(f),
414 }
415 }
416
417 pub(crate) fn try_for_each<F, E>(self, mut f: F) -> Result<(), E>
420 where
421 F: FnMut(M::Item) -> Result<(), E>,
422 {
423 match self {
424 Self::Materialized(iter) => {
425 for item in iter {
426 f(item)?;
427 }
428 }
429 Self::Lazy(iter) => {
430 for item in iter {
431 f(item)?;
432 }
433 }
434 }
435
436 Ok(())
437 }
438}
439
440#[derive(Debug)]
442pub struct FilterPredicate {
443 filter: BooleanArray,
444 count: usize,
445 strategy: IterationStrategy,
447}
448
449impl FilterPredicate {
450 pub fn filter(&self, values: &dyn Array) -> Result<ArrayRef, ArrowError> {
452 filter_array(values, self)
453 }
454
455 pub fn filter_record_batch(
460 &self,
461 record_batch: &RecordBatch,
462 ) -> Result<RecordBatch, ArrowError> {
463 let filtered_arrays = record_batch
464 .columns()
465 .iter()
466 .map(|a| filter_array(a, self))
467 .collect::<Result<Vec<_>, _>>()?;
468
469 unsafe {
472 Ok(RecordBatch::new_unchecked(
473 record_batch.schema(),
474 filtered_arrays,
475 self.count,
476 ))
477 }
478 }
479
480 pub fn count(&self) -> usize {
482 self.count
483 }
484
485 pub(crate) fn selection(&self) -> FilterSelection<'_> {
488 match &self.strategy {
489 IterationStrategy::None => FilterSelection::None,
490 IterationStrategy::All => FilterSelection::All { len: self.count },
491 IterationStrategy::Slices(slices) => {
492 FilterSelection::Slices(FilterIterator::Materialized(slices.iter().copied()))
493 }
494 IterationStrategy::SlicesIterator => {
495 FilterSelection::Slices(FilterIterator::Lazy(SlicesIterator::new(&self.filter)))
496 }
497 IterationStrategy::Indices(indices) => {
498 FilterSelection::Indices(FilterIterator::Materialized(indices.iter().copied()))
499 }
500 IterationStrategy::IndexIterator => FilterSelection::Indices(FilterIterator::Lazy(
501 IndexIterator::new(&self.filter, self.count),
502 )),
503 }
504 }
505
506 pub fn filter_nulls(&self, nulls: Option<&NullBuffer>) -> Option<NullBuffer> {
513 let nulls = nulls?;
514 if nulls.null_count() == 0 {
515 return None;
516 }
517
518 let nulls = filter_bits(nulls.inner(), self);
519 let null_count = self.count - nulls.count_set_bits_offset(0, self.count);
522
523 if null_count == 0 {
524 return None;
525 }
526
527 let buffer = BooleanBuffer::new(nulls, 0, self.count);
528 debug_assert_eq!(null_count, buffer.len() - buffer.count_set_bits());
529 Some(unsafe { NullBuffer::new_unchecked(buffer, null_count) })
532 }
533}
534
535fn filter_array(values: &dyn Array, predicate: &FilterPredicate) -> Result<ArrayRef, ArrowError> {
536 if predicate.filter.len() > values.len() {
537 return Err(ArrowError::InvalidArgumentError(format!(
538 "Filter predicate of length {} is larger than target array of length {}",
539 predicate.filter.len(),
540 values.len()
541 )));
542 }
543
544 match predicate.strategy {
545 IterationStrategy::None => Ok(new_empty_array(values.data_type())),
546 IterationStrategy::All => Ok(values.slice(0, predicate.count)),
547 _ => downcast_primitive_array! {
549 values => Ok(Arc::new(filter_primitive(values, predicate))),
550 DataType::Boolean => {
551 let values = values.as_any().downcast_ref::<BooleanArray>().unwrap();
552 Ok(Arc::new(filter_boolean(values, predicate)))
553 }
554 DataType::Utf8 => {
555 Ok(Arc::new(filter_bytes(values.as_string::<i32>(), predicate)))
556 }
557 DataType::LargeUtf8 => {
558 Ok(Arc::new(filter_bytes(values.as_string::<i64>(), predicate)))
559 }
560 DataType::Utf8View => {
561 Ok(Arc::new(filter_byte_view(values.as_string_view(), predicate)))
562 }
563 DataType::Binary => {
564 Ok(Arc::new(filter_bytes(values.as_binary::<i32>(), predicate)))
565 }
566 DataType::LargeBinary => {
567 Ok(Arc::new(filter_bytes(values.as_binary::<i64>(), predicate)))
568 }
569 DataType::BinaryView => {
570 Ok(Arc::new(filter_byte_view(values.as_binary_view(), predicate)))
571 }
572 DataType::FixedSizeBinary(_) => {
573 Ok(Arc::new(filter_fixed_size_binary(values.as_fixed_size_binary(), predicate)))
574 }
575 DataType::ListView(_) => {
576 Ok(Arc::new(filter_list_view::<i32>(values.as_list_view(), predicate)))
577 }
578 DataType::LargeListView(_) => {
579 Ok(Arc::new(filter_list_view::<i64>(values.as_list_view(), predicate)))
580 }
581 DataType::RunEndEncoded(_, _) => {
582 downcast_run_array!{
583 values => Ok(Arc::new(filter_run_end_array(values, predicate)?)),
584 t => unimplemented!("Filter not supported for RunEndEncoded type {:?}", t)
585 }
586 }
587 DataType::Dictionary(_, _) => downcast_dictionary_array! {
588 values => Ok(Arc::new(filter_dict(values, predicate))),
589 t => unimplemented!("Filter not supported for dictionary type {:?}", t)
590 }
591 DataType::Struct(_) => {
592 Ok(Arc::new(filter_struct(values.as_struct(), predicate)?))
593 }
594 DataType::Union(_, UnionMode::Sparse) => {
595 Ok(Arc::new(filter_sparse_union(values.as_union(), predicate)?))
596 }
597 _ => {
598 let data = values.to_data();
599 let mut mutable = MutableArrayData::new(
601 vec![&data],
602 false,
603 predicate.count,
604 );
605
606 match &predicate.strategy {
607 IterationStrategy::Slices(slices) => {
608 for (start, end) in slices {
609 mutable.try_extend(0, *start, *end)?;
610 }
611 }
612 _ => {
613 let iter = SlicesIterator::new(&predicate.filter);
614 for (start, end) in iter {
615 mutable.try_extend(0, start, end)?;
616 }
617 }
618 }
619
620 let data = mutable.freeze();
621 Ok(make_array(data))
622 }
623 },
624 }
625}
626
627fn filter_run_end_array<R: RunEndIndexType>(
629 array: &RunArray<R>,
630 predicate: &FilterPredicate,
631) -> Result<RunArray<R>, ArrowError>
632where
633 R::Native: Into<i64> + From<bool>,
634 R::Native: AddAssign,
635{
636 let run_ends: &RunEndBuffer<R::Native> = array.run_ends();
637 let start_physical = run_ends.get_start_physical_index();
638 let end_physical = run_ends.get_end_physical_index();
639 let physical_len = end_physical - start_physical + 1;
640
641 let mut new_run_ends = vec![R::default_value(); physical_len];
642 let offset = run_ends.offset() as u64;
643
644 let mut start = 0u64;
645 let mut j = 0;
646 let mut count = R::default_value();
647 let filter_values = predicate.filter.values();
648 let run_ends = run_ends.inner();
649
650 let pred: BooleanArray = BooleanBuffer::collect_bool(physical_len, |i| {
651 let mut keep = false;
652 let mut end = (run_ends[i + start_physical].into() as u64).saturating_sub(offset);
653 let difference = end.saturating_sub(filter_values.len() as u64);
654 end -= difference;
655
656 for pred in (start..end).map(|i| unsafe { filter_values.value_unchecked(i as usize) }) {
658 count += R::Native::from(pred);
659 keep |= pred
660 }
661 new_run_ends[j] = count;
663 j += keep as usize;
664
665 start = end;
666 keep
667 })
668 .into();
669
670 new_run_ends.truncate(j);
671
672 let values = array.values_slice();
673 let values = filter(values.as_ref(), &pred)?;
674
675 let run_ends = PrimitiveArray::<R>::try_new(new_run_ends.into(), None)?;
676 RunArray::try_new(&run_ends, &values)
677}
678
679fn filter_bits(buffer: &BooleanBuffer, predicate: &FilterPredicate) -> Buffer {
681 let src = buffer.values();
682 let offset = buffer.offset();
683 assert!(buffer.len() >= predicate.filter.len());
684
685 match &predicate.strategy {
686 IterationStrategy::IndexIterator => {
687 let bits =
688 IndexIterator::new(&predicate.filter, predicate.count).map(|src_idx| unsafe {
690 bit_util::get_bit_raw(buffer.values().as_ptr(), src_idx + offset)
691 });
692
693 unsafe { MutableBuffer::from_trusted_len_iter_bool(bits).into() }
695 }
696 IterationStrategy::Indices(indices) => {
697 let bits = indices.iter().map(|src_idx| unsafe {
699 bit_util::get_bit_raw(buffer.values().as_ptr(), *src_idx + offset)
700 });
701 unsafe { MutableBuffer::from_trusted_len_iter_bool(bits).into() }
703 }
704 IterationStrategy::SlicesIterator => {
705 let mut builder = BooleanBufferBuilder::new(predicate.count);
706 for (start, end) in SlicesIterator::new(&predicate.filter) {
707 builder.append_packed_range(start + offset..end + offset, src)
708 }
709 builder.into()
710 }
711 IterationStrategy::Slices(slices) => {
712 let mut builder = BooleanBufferBuilder::new(predicate.count);
713 for (start, end) in slices {
714 builder.append_packed_range(*start + offset..*end + offset, src)
715 }
716 builder.into()
717 }
718 IterationStrategy::All | IterationStrategy::None => unreachable!(),
719 }
720}
721
722fn filter_boolean(array: &BooleanArray, predicate: &FilterPredicate) -> BooleanArray {
724 let buffer = filter_bits(array.values(), predicate);
725 let values = BooleanBuffer::new(buffer, 0, predicate.count);
726 let nulls = predicate.filter_nulls(array.nulls());
727
728 BooleanArray::new(values, nulls)
729}
730
731#[inline(never)]
732pub(crate) fn filter_native<T: ArrowNativeType>(
733 values: &[T],
734 predicate: &FilterPredicate,
735) -> Buffer {
736 assert!(values.len() >= predicate.filter.len());
737
738 match &predicate.strategy {
739 IterationStrategy::SlicesIterator => {
740 let mut buffer = Vec::with_capacity(predicate.count);
741 for (start, end) in SlicesIterator::new(&predicate.filter) {
742 buffer.extend_from_slice(unsafe { values.get_unchecked(start..end) });
744 }
745 buffer.into()
746 }
747 IterationStrategy::Slices(slices) => {
748 let mut buffer = Vec::with_capacity(predicate.count);
749 for (start, end) in slices {
750 buffer.extend_from_slice(unsafe { values.get_unchecked(*start..*end) });
752 }
753 buffer.into()
754 }
755 IterationStrategy::IndexIterator => {
756 let iter = IndexIterator::new(&predicate.filter, predicate.count)
758 .map(|x| unsafe { *values.get_unchecked(x) });
759
760 unsafe { MutableBuffer::from_trusted_len_iter(iter) }.into()
762 }
763 IterationStrategy::Indices(indices) => {
764 let iter = indices.iter().map(|x| unsafe { *values.get_unchecked(*x) });
766 iter.collect::<Vec<_>>().into()
767 }
768 IterationStrategy::All | IterationStrategy::None => unreachable!(),
769 }
770}
771
772fn filter_primitive<T>(array: &PrimitiveArray<T>, predicate: &FilterPredicate) -> PrimitiveArray<T>
774where
775 T: ArrowPrimitiveType,
776{
777 let buffer = filter_native(array.values(), predicate);
778 let values = ScalarBuffer::new(buffer, 0, predicate.count);
779 let nulls = predicate.filter_nulls(array.nulls());
780 let filtered = PrimitiveArray::new(values, nulls);
781
782 if array.data_type() == &T::DATA_TYPE {
784 filtered
785 } else {
786 filtered.with_data_type(array.data_type().clone())
787 }
788}
789
790struct FilterBytes<'a, OffsetSize> {
795 src_offsets: &'a [OffsetSize],
796 src_values: &'a [u8],
797 dst_offsets: Vec<OffsetSize>,
798 dst_values: Vec<u8>,
799 cur_offset: OffsetSize,
800}
801
802impl<'a, OffsetSize> FilterBytes<'a, OffsetSize>
803where
804 OffsetSize: OffsetSizeTrait,
805{
806 fn new<T>(capacity: usize, array: &'a GenericByteArray<T>) -> Self
807 where
808 T: ByteArrayType<Offset = OffsetSize>,
809 {
810 let dst_values = Vec::new();
811 let mut dst_offsets: Vec<OffsetSize> = Vec::with_capacity(capacity + 1);
812 let cur_offset = OffsetSize::from_usize(0).unwrap();
813
814 dst_offsets.push(cur_offset);
815
816 Self {
817 src_offsets: array.value_offsets(),
818 src_values: array.value_data(),
819 dst_offsets,
820 dst_values,
821 cur_offset,
822 }
823 }
824
825 #[inline]
827 fn get_value_offset(&self, idx: usize) -> usize {
828 self.src_offsets[idx].as_usize()
829 }
830
831 #[inline]
833 fn get_value_range(&self, idx: usize) -> (usize, usize, OffsetSize) {
834 let start = self.get_value_offset(idx);
836 let end = self.get_value_offset(idx + 1);
837 let len = OffsetSize::from_usize(end - start).expect("illegal offset range");
838 (start, end, len)
839 }
840
841 fn extend_offsets_idx(&mut self, iter: impl Iterator<Item = usize>) {
842 self.dst_offsets.extend(iter.map(|idx| {
843 let start = self.src_offsets[idx].as_usize();
844 let end = self.src_offsets[idx + 1].as_usize();
845 let len = OffsetSize::from_usize(end - start).expect("illegal offset range");
846 self.cur_offset += len;
847
848 self.cur_offset
849 }));
850 }
851
852 fn extend_idx(&mut self, iter: impl Iterator<Item = usize>) {
854 self.dst_values.reserve_exact(self.cur_offset.as_usize());
855
856 for idx in iter {
857 let start = self.src_offsets[idx].as_usize();
858 let end = self.src_offsets[idx + 1].as_usize();
859 self.dst_values
860 .extend_from_slice(&self.src_values[start..end]);
861 }
862 }
863
864 fn extend_offsets_slices(&mut self, iter: impl Iterator<Item = (usize, usize)>, count: usize) {
865 self.dst_offsets.reserve_exact(count);
866 for (start, end) in iter {
867 for idx in start..end {
869 let (_, _, len) = self.get_value_range(idx);
870 self.cur_offset += len;
871 self.dst_offsets.push(self.cur_offset);
872 }
873 }
874 }
875
876 fn extend_slices(&mut self, iter: impl Iterator<Item = (usize, usize)>) {
878 self.dst_values.reserve_exact(self.cur_offset.as_usize());
879
880 for (start, end) in iter {
881 let value_start = self.get_value_offset(start);
882 let value_end = self.get_value_offset(end);
883 self.dst_values
884 .extend_from_slice(&self.src_values[value_start..value_end]);
885 }
886 }
887}
888
889fn filter_bytes<T>(array: &GenericByteArray<T>, predicate: &FilterPredicate) -> GenericByteArray<T>
894where
895 T: ByteArrayType,
896{
897 let mut filter = FilterBytes::new(predicate.count, array);
898
899 match &predicate.strategy {
900 IterationStrategy::SlicesIterator => {
901 filter.extend_offsets_slices(SlicesIterator::new(&predicate.filter), predicate.count);
902 filter.extend_slices(SlicesIterator::new(&predicate.filter))
903 }
904 IterationStrategy::Slices(slices) => {
905 filter.extend_offsets_slices(slices.iter().copied(), predicate.count);
906 filter.extend_slices(slices.iter().copied())
907 }
908 IterationStrategy::IndexIterator => {
909 filter.extend_offsets_idx(IndexIterator::new(&predicate.filter, predicate.count));
910 filter.extend_idx(IndexIterator::new(&predicate.filter, predicate.count))
911 }
912 IterationStrategy::Indices(indices) => {
913 filter.extend_offsets_idx(indices.iter().copied());
914 filter.extend_idx(indices.iter().copied())
915 }
916 IterationStrategy::All | IterationStrategy::None => unreachable!(),
917 }
918
919 let offsets = unsafe { OffsetBuffer::new_unchecked(filter.dst_offsets.into()) };
922 let nulls = predicate.filter_nulls(array.nulls());
923
924 unsafe { GenericByteArray::new_unchecked(offsets, filter.dst_values.into(), nulls) }
928}
929
930fn filter_byte_view<T: ByteViewType>(
932 array: &GenericByteViewArray<T>,
933 predicate: &FilterPredicate,
934) -> GenericByteViewArray<T> {
935 let new_view_buffer = filter_native(array.views(), predicate);
936 let views = ScalarBuffer::new(new_view_buffer, 0, predicate.count);
937 let buffers = Arc::clone(array.data_buffers());
938 let nulls = predicate.filter_nulls(array.nulls());
939
940 unsafe { GenericByteViewArray::new_unchecked(views, buffers, nulls) }
944}
945
946#[inline(always)]
949fn copy_fsb_indices(
950 values: &[u8],
951 value_length: usize,
952 indices: impl Iterator<Item = usize>,
953 count: usize,
954) -> MutableBuffer {
955 let total = count * value_length;
956 let mut buffer = MutableBuffer::with_capacity(total);
957 let dst_base = buffer.as_mut_ptr();
958 let mut write_offset = 0usize;
959 for idx in indices {
960 let src_start = idx * value_length;
961 unsafe {
964 std::ptr::copy_nonoverlapping(
965 values.as_ptr().add(src_start),
966 dst_base.add(write_offset),
967 value_length,
968 );
969 }
970 write_offset += value_length;
971 }
972 unsafe { buffer.set_len(total) };
974 buffer
975}
976
977fn filter_fixed_size_binary(
978 array: &FixedSizeBinaryArray,
979 predicate: &FilterPredicate,
980) -> FixedSizeBinaryArray {
981 let values: &[u8] = array.values();
982 let value_length = array.value_length() as usize;
983 let calculate_offset_from_index = |index: usize| index * value_length;
984 let buffer = match &predicate.strategy {
985 IterationStrategy::SlicesIterator => {
986 let mut buffer = MutableBuffer::with_capacity(predicate.count * value_length);
987 for (start, end) in SlicesIterator::new(&predicate.filter) {
988 buffer.extend_from_slice(
989 &values[calculate_offset_from_index(start)..calculate_offset_from_index(end)],
990 );
991 }
992 buffer
993 }
994 IterationStrategy::Slices(slices) => {
995 let mut buffer = MutableBuffer::with_capacity(predicate.count * value_length);
996 for (start, end) in slices {
997 buffer.extend_from_slice(
998 &values[calculate_offset_from_index(*start)..calculate_offset_from_index(*end)],
999 );
1000 }
1001 buffer
1002 }
1003 IterationStrategy::IndexIterator => copy_fsb_indices(
1004 values,
1005 value_length,
1006 IndexIterator::new(&predicate.filter, predicate.count),
1007 predicate.count,
1008 ),
1009 IterationStrategy::Indices(indices) => copy_fsb_indices(
1010 values,
1011 value_length,
1012 indices.iter().copied(),
1013 predicate.count,
1014 ),
1015 IterationStrategy::All | IterationStrategy::None => unreachable!(),
1016 };
1017
1018 let nulls = predicate.filter_nulls(array.nulls());
1019
1020 FixedSizeBinaryArray::new(array.value_length(), buffer.into(), nulls)
1021}
1022
1023fn filter_dict<K: ArrowDictionaryKeyType>(
1025 array: &DictionaryArray<K>,
1026 predicate: &FilterPredicate,
1027) -> DictionaryArray<K> {
1028 let new_keys = filter_primitive(array.keys(), predicate);
1031 unsafe { DictionaryArray::new_unchecked(new_keys, array.values().clone()) }
1032}
1033
1034fn filter_struct(
1036 array: &StructArray,
1037 predicate: &FilterPredicate,
1038) -> Result<StructArray, ArrowError> {
1039 let columns = array
1040 .columns()
1041 .iter()
1042 .map(|column| filter_array(column, predicate))
1043 .collect::<Result<_, _>>()?;
1044
1045 let nulls = predicate.filter_nulls(array.nulls());
1046
1047 Ok(unsafe {
1048 StructArray::new_unchecked_with_length(
1049 array.fields().clone(),
1050 columns,
1051 nulls,
1052 predicate.count(),
1053 )
1054 })
1055}
1056
1057fn filter_sparse_union(
1059 array: &UnionArray,
1060 predicate: &FilterPredicate,
1061) -> Result<UnionArray, ArrowError> {
1062 let DataType::Union(fields, UnionMode::Sparse) = array.data_type() else {
1063 unreachable!()
1064 };
1065
1066 let type_ids = filter_primitive(
1067 &Int8Array::try_new(array.type_ids().clone(), None)?,
1068 predicate,
1069 );
1070
1071 let children = fields
1072 .iter()
1073 .map(|(child_type_id, _)| filter_array(array.child(child_type_id), predicate))
1074 .collect::<Result<_, _>>()?;
1075
1076 Ok(unsafe {
1077 UnionArray::new_unchecked(fields.clone(), type_ids.into_parts().1, None, children)
1078 })
1079}
1080
1081fn filter_list_view<OffsetType: OffsetSizeTrait>(
1083 array: &GenericListViewArray<OffsetType>,
1084 predicate: &FilterPredicate,
1085) -> GenericListViewArray<OffsetType> {
1086 let filtered_offsets = filter_native::<OffsetType>(array.offsets(), predicate);
1087 let filtered_sizes = filter_native::<OffsetType>(array.sizes(), predicate);
1088
1089 let field = match array.data_type() {
1090 DataType::ListView(field) | DataType::LargeListView(field) => field.clone(),
1091 _ => unreachable!(),
1092 };
1093 let offsets = ScalarBuffer::new(filtered_offsets, 0, predicate.count);
1094 let sizes = ScalarBuffer::new(filtered_sizes, 0, predicate.count);
1095 let values = array.values().clone();
1096 let nulls = predicate.filter_nulls(array.nulls());
1097
1098 unsafe { GenericListViewArray::new_unchecked(field, offsets, sizes, values, nulls) }
1102}
1103
1104#[cfg(test)]
1105mod tests {
1106 use super::*;
1107 use arrow_array::builder::*;
1108 use arrow_array::cast::as_run_array;
1109 use arrow_array::types::*;
1110 use rand::distr::uniform::{UniformSampler, UniformUsize};
1111 use rand::distr::{Alphanumeric, StandardUniform};
1112 use rand::prelude::*;
1113 use rand::rng;
1114
1115 macro_rules! def_temporal_test {
1116 ($test:ident, $array_type: ident, $data: expr) => {
1117 #[test]
1118 fn $test() {
1119 let a = $data;
1120 let b = BooleanArray::from(vec![true, false, true, false]);
1121 let c = filter(&a, &b).unwrap();
1122 let d = c.as_ref().as_any().downcast_ref::<$array_type>().unwrap();
1123 assert_eq!(2, d.len());
1124 assert_eq!(1, d.value(0));
1125 assert_eq!(3, d.value(1));
1126 }
1127 };
1128 }
1129
1130 def_temporal_test!(
1131 test_filter_date32,
1132 Date32Array,
1133 Date32Array::from(vec![1, 2, 3, 4])
1134 );
1135 def_temporal_test!(
1136 test_filter_date64,
1137 Date64Array,
1138 Date64Array::from(vec![1, 2, 3, 4])
1139 );
1140 def_temporal_test!(
1141 test_filter_time32_second,
1142 Time32SecondArray,
1143 Time32SecondArray::from(vec![1, 2, 3, 4])
1144 );
1145 def_temporal_test!(
1146 test_filter_time32_millisecond,
1147 Time32MillisecondArray,
1148 Time32MillisecondArray::from(vec![1, 2, 3, 4])
1149 );
1150 def_temporal_test!(
1151 test_filter_time64_microsecond,
1152 Time64MicrosecondArray,
1153 Time64MicrosecondArray::from(vec![1, 2, 3, 4])
1154 );
1155 def_temporal_test!(
1156 test_filter_time64_nanosecond,
1157 Time64NanosecondArray,
1158 Time64NanosecondArray::from(vec![1, 2, 3, 4])
1159 );
1160 def_temporal_test!(
1161 test_filter_duration_second,
1162 DurationSecondArray,
1163 DurationSecondArray::from(vec![1, 2, 3, 4])
1164 );
1165 def_temporal_test!(
1166 test_filter_duration_millisecond,
1167 DurationMillisecondArray,
1168 DurationMillisecondArray::from(vec![1, 2, 3, 4])
1169 );
1170 def_temporal_test!(
1171 test_filter_duration_microsecond,
1172 DurationMicrosecondArray,
1173 DurationMicrosecondArray::from(vec![1, 2, 3, 4])
1174 );
1175 def_temporal_test!(
1176 test_filter_duration_nanosecond,
1177 DurationNanosecondArray,
1178 DurationNanosecondArray::from(vec![1, 2, 3, 4])
1179 );
1180 def_temporal_test!(
1181 test_filter_timestamp_second,
1182 TimestampSecondArray,
1183 TimestampSecondArray::from(vec![1, 2, 3, 4])
1184 );
1185 def_temporal_test!(
1186 test_filter_timestamp_millisecond,
1187 TimestampMillisecondArray,
1188 TimestampMillisecondArray::from(vec![1, 2, 3, 4])
1189 );
1190 def_temporal_test!(
1191 test_filter_timestamp_microsecond,
1192 TimestampMicrosecondArray,
1193 TimestampMicrosecondArray::from(vec![1, 2, 3, 4])
1194 );
1195 def_temporal_test!(
1196 test_filter_timestamp_nanosecond,
1197 TimestampNanosecondArray,
1198 TimestampNanosecondArray::from(vec![1, 2, 3, 4])
1199 );
1200
1201 #[test]
1202 fn test_filter_array_slice() {
1203 let a = Int32Array::from(vec![5, 6, 7, 8, 9]).slice(1, 4);
1204 let b = BooleanArray::from(vec![true, false, false, true]);
1205 let c = filter(&a, &b).unwrap();
1209 let d = c.as_ref().as_any().downcast_ref::<Int32Array>().unwrap();
1210 assert_eq!(2, d.len());
1211 assert_eq!(6, d.value(0));
1212 assert_eq!(9, d.value(1));
1213 }
1214
1215 #[test]
1216 fn test_filter_array_low_density() {
1217 let mut data_values = (1..=65).collect::<Vec<i32>>();
1219 let mut filter_values = (1..=65).map(|i| matches!(i % 65, 0)).collect::<Vec<bool>>();
1220 data_values.extend_from_slice(&[66, 67]);
1222 filter_values.extend_from_slice(&[false, true]);
1223 let a = Int32Array::from(data_values);
1224 let b = BooleanArray::from(filter_values);
1225 let c = filter(&a, &b).unwrap();
1226 let d = c.as_ref().as_any().downcast_ref::<Int32Array>().unwrap();
1227 assert_eq!(2, d.len());
1228 assert_eq!(65, d.value(0));
1229 assert_eq!(67, d.value(1));
1230 }
1231
1232 #[test]
1233 fn test_filter_array_high_density() {
1234 let mut data_values = (1..=65).map(Some).collect::<Vec<_>>();
1236 let mut filter_values = (1..=65)
1237 .map(|i| !matches!(i % 65, 0))
1238 .collect::<Vec<bool>>();
1239 data_values[1] = None;
1241 data_values.extend_from_slice(&[Some(66), None, Some(67), None]);
1243 filter_values.extend_from_slice(&[false, true, true, true]);
1244 let a = Int32Array::from(data_values);
1245 let b = BooleanArray::from(filter_values);
1246 let c = filter(&a, &b).unwrap();
1247 let d = c.as_ref().as_any().downcast_ref::<Int32Array>().unwrap();
1248 assert_eq!(67, d.len());
1249 assert_eq!(3, d.null_count());
1250 assert_eq!(1, d.value(0));
1251 assert!(d.is_null(1));
1252 assert_eq!(64, d.value(63));
1253 assert!(d.is_null(64));
1254 assert_eq!(67, d.value(65));
1255 }
1256
1257 #[test]
1258 fn test_filter_string_array_simple() {
1259 let a = StringArray::from(vec!["hello", " ", "world", "!"]);
1260 let b = BooleanArray::from(vec![true, false, true, false]);
1261 let c = filter(&a, &b).unwrap();
1262 let d = c.as_ref().as_any().downcast_ref::<StringArray>().unwrap();
1263 assert_eq!(2, d.len());
1264 assert_eq!("hello", d.value(0));
1265 assert_eq!("world", d.value(1));
1266 }
1267
1268 #[test]
1269 fn test_filter_primitive_array_with_null() {
1270 let a = Int32Array::from(vec![Some(5), None]);
1271 let b = BooleanArray::from(vec![false, true]);
1272 let c = filter(&a, &b).unwrap();
1273 let d = c.as_ref().as_any().downcast_ref::<Int32Array>().unwrap();
1274 assert_eq!(1, d.len());
1275 assert!(d.is_null(0));
1276 }
1277
1278 #[test]
1279 fn test_filter_string_array_with_null() {
1280 let a = StringArray::from(vec![Some("hello"), None, Some("world"), None]);
1281 let b = BooleanArray::from(vec![true, false, false, true]);
1282 let c = filter(&a, &b).unwrap();
1283 let d = c.as_ref().as_any().downcast_ref::<StringArray>().unwrap();
1284 assert_eq!(2, d.len());
1285 assert_eq!("hello", d.value(0));
1286 assert!(!d.is_null(0));
1287 assert!(d.is_null(1));
1288 }
1289
1290 #[test]
1291 fn test_filter_binary_array_with_null() {
1292 let data: Vec<Option<&[u8]>> = vec![Some(b"hello"), None, Some(b"world"), None];
1293 let a = BinaryArray::from(data);
1294 let b = BooleanArray::from(vec![true, false, false, true]);
1295 let c = filter(&a, &b).unwrap();
1296 let d = c.as_ref().as_any().downcast_ref::<BinaryArray>().unwrap();
1297 assert_eq!(2, d.len());
1298 assert_eq!(b"hello", d.value(0));
1299 assert!(!d.is_null(0));
1300 assert!(d.is_null(1));
1301 }
1302
1303 fn _test_filter_byte_view<T>()
1304 where
1305 T: ByteViewType,
1306 str: AsRef<T::Native>,
1307 T::Native: PartialEq,
1308 {
1309 let array = {
1310 let mut builder = GenericByteViewBuilder::<T>::new();
1312 builder.append_value("hello");
1313 builder.append_value("world");
1314 builder.append_null();
1315 builder.append_value("large payload over 12 bytes");
1316 builder.append_value("lulu");
1317 builder.finish()
1318 };
1319
1320 {
1321 let predicate = BooleanArray::from(vec![true, false, true, true, false]);
1322 let actual = filter(&array, &predicate).unwrap();
1323
1324 assert_eq!(actual.len(), 3);
1325 let actual_buffers = actual.as_byte_view::<T>().data_buffers();
1326 let input_buffers = array.data_buffers();
1327 assert!(Arc::ptr_eq(actual_buffers, input_buffers));
1328
1329 let expected = {
1330 let mut builder = GenericByteViewBuilder::<T>::new();
1332 builder.append_value("hello");
1333 builder.append_null();
1334 builder.append_value("large payload over 12 bytes");
1335 builder.finish()
1336 };
1337
1338 assert_eq!(actual.as_ref(), &expected);
1339 }
1340
1341 {
1342 let predicate = BooleanArray::from(vec![true, false, false, false, true]);
1343 let actual = filter(&array, &predicate).unwrap();
1344
1345 assert_eq!(actual.len(), 2);
1346
1347 let expected = {
1348 let mut builder = GenericByteViewBuilder::<T>::new();
1350 builder.append_value("hello");
1351 builder.append_value("lulu");
1352 builder.finish()
1353 };
1354
1355 assert_eq!(actual.as_ref(), &expected);
1356 }
1357 }
1358
1359 #[test]
1360 fn test_filter_string_view() {
1361 _test_filter_byte_view::<StringViewType>()
1362 }
1363
1364 #[test]
1365 fn test_filter_binary_view() {
1366 _test_filter_byte_view::<BinaryViewType>()
1367 }
1368
1369 #[test]
1370 fn test_filter_fixed_binary() {
1371 let v1 = [1_u8, 2];
1372 let v2 = [3_u8, 4];
1373 let v3 = [5_u8, 6];
1374 let v = vec![&v1, &v2, &v3];
1375 let a = FixedSizeBinaryArray::try_from(v).unwrap();
1376 let b = BooleanArray::from(vec![true, false, true]);
1377 let c = filter(&a, &b).unwrap();
1378 let d = c
1379 .as_ref()
1380 .as_any()
1381 .downcast_ref::<FixedSizeBinaryArray>()
1382 .unwrap();
1383 assert_eq!(d.len(), 2);
1384 assert_eq!(d.value(0), &v1);
1385 assert_eq!(d.value(1), &v3);
1386 let c2 = FilterBuilder::new(&b)
1387 .optimize()
1388 .build()
1389 .filter(&a)
1390 .unwrap();
1391 let d2 = c2
1392 .as_ref()
1393 .as_any()
1394 .downcast_ref::<FixedSizeBinaryArray>()
1395 .unwrap();
1396 assert_eq!(d, d2);
1397
1398 let b = BooleanArray::from(vec![false, false, false]);
1399 let c = filter(&a, &b).unwrap();
1400 let d = c
1401 .as_ref()
1402 .as_any()
1403 .downcast_ref::<FixedSizeBinaryArray>()
1404 .unwrap();
1405 assert_eq!(d.len(), 0);
1406
1407 let b = BooleanArray::from(vec![true, true, true]);
1408 let c = filter(&a, &b).unwrap();
1409 let d = c
1410 .as_ref()
1411 .as_any()
1412 .downcast_ref::<FixedSizeBinaryArray>()
1413 .unwrap();
1414 assert_eq!(d.len(), 3);
1415 assert_eq!(d.value(0), &v1);
1416 assert_eq!(d.value(1), &v2);
1417 assert_eq!(d.value(2), &v3);
1418
1419 let b = BooleanArray::from(vec![false, false, true]);
1420 let c = filter(&a, &b).unwrap();
1421 let d = c
1422 .as_ref()
1423 .as_any()
1424 .downcast_ref::<FixedSizeBinaryArray>()
1425 .unwrap();
1426 assert_eq!(d.len(), 1);
1427 assert_eq!(d.value(0), &v3);
1428 let c2 = FilterBuilder::new(&b)
1429 .optimize()
1430 .build()
1431 .filter(&a)
1432 .unwrap();
1433 let d2 = c2
1434 .as_ref()
1435 .as_any()
1436 .downcast_ref::<FixedSizeBinaryArray>()
1437 .unwrap();
1438 assert_eq!(d, d2);
1439 }
1440
1441 #[test]
1442 fn test_filter_array_slice_with_null() {
1443 let a = Int32Array::from(vec![Some(5), None, Some(7), Some(8), Some(9)]).slice(1, 4);
1444 let b = BooleanArray::from(vec![true, false, false, true]);
1445 let c = filter(&a, &b).unwrap();
1449 let d = c.as_ref().as_any().downcast_ref::<Int32Array>().unwrap();
1450 assert_eq!(2, d.len());
1451 assert!(d.is_null(0));
1452 assert!(!d.is_null(1));
1453 assert_eq!(9, d.value(1));
1454 }
1455
1456 #[test]
1457 fn test_filter_run_end_encoding_array() {
1458 let run_ends = Int64Array::from(vec![2, 3, 8]);
1459 let values = Int64Array::from(vec![7, -2, 9]);
1460 let a = RunArray::try_new(&run_ends, &values).expect("Failed to create RunArray");
1461 let b = BooleanArray::from(vec![true, false, true, false, true, false, true, false]);
1462 let c = filter(&a, &b).unwrap();
1463 let actual: &RunArray<Int64Type> = as_run_array(&c);
1464 assert_eq!(4, actual.len());
1465
1466 let expected = RunArray::try_new(
1467 &Int64Array::from(vec![1, 2, 4]),
1468 &Int64Array::from(vec![7, -2, 9]),
1469 )
1470 .expect("Failed to make expected RunArray test is broken");
1471
1472 assert_eq!(&actual.run_ends().values(), &expected.run_ends().values());
1473 assert_eq!(actual.values(), expected.values())
1474 }
1475
1476 #[test]
1477 fn test_filter_run_end_encoding_array_sliced() {
1478 let run_ends = Int64Array::from(vec![2, 3, 8]);
1479 let values = Int64Array::from(vec![7, -2, 9]);
1480 let a = RunArray::try_new(&run_ends, &values).unwrap(); let a = a.slice(2, 3); let b = BooleanArray::from(vec![true, false, true]);
1483 let result = filter(&a, &b).unwrap();
1484
1485 let result = result.as_run::<Int64Type>();
1486 let result = result.downcast::<Int64Array>().unwrap();
1487
1488 let expected = vec![-2, 9];
1489 let actual = result.into_iter().flatten().collect::<Vec<_>>();
1490 assert_eq!(expected, actual);
1491 }
1492
1493 #[test]
1494 fn test_filter_run_end_encoding_array_remove_value() {
1495 let run_ends = Int32Array::from(vec![2, 3, 8, 10]);
1496 let values = Int32Array::from(vec![7, -2, 9, -8]);
1497 let a = RunArray::try_new(&run_ends, &values).expect("Failed to create RunArray");
1498 let b = BooleanArray::from(vec![
1499 false, true, false, false, true, false, true, false, false, false,
1500 ]);
1501 let c = filter(&a, &b).unwrap();
1502 let actual: &RunArray<Int32Type> = as_run_array(&c);
1503 assert_eq!(3, actual.len());
1504
1505 let expected =
1506 RunArray::try_new(&Int32Array::from(vec![1, 3]), &Int32Array::from(vec![7, 9]))
1507 .expect("Failed to make expected RunArray test is broken");
1508
1509 assert_eq!(&actual.run_ends().values(), &expected.run_ends().values());
1510 assert_eq!(actual.values(), expected.values())
1511 }
1512
1513 #[test]
1514 fn test_filter_run_end_encoding_array_remove_all_but_one() {
1515 let run_ends = Int16Array::from(vec![2, 3, 8, 10]);
1516 let values = Int16Array::from(vec![7, -2, 9, -8]);
1517 let a = RunArray::try_new(&run_ends, &values).expect("Failed to create RunArray");
1518 let b = BooleanArray::from(vec![
1519 false, false, false, false, false, false, true, false, false, false,
1520 ]);
1521 let c = filter(&a, &b).unwrap();
1522 let actual: &RunArray<Int16Type> = as_run_array(&c);
1523 assert_eq!(1, actual.len());
1524
1525 let expected = RunArray::try_new(&Int16Array::from(vec![1]), &Int16Array::from(vec![9]))
1526 .expect("Failed to make expected RunArray test is broken");
1527
1528 assert_eq!(&actual.run_ends().values(), &expected.run_ends().values());
1529 assert_eq!(actual.values(), expected.values())
1530 }
1531
1532 #[test]
1533 fn test_filter_run_end_encoding_array_empty() {
1534 let run_ends = Int64Array::from(vec![2, 3, 8, 10]);
1535 let values = Int64Array::from(vec![7, -2, 9, -8]);
1536 let a = RunArray::try_new(&run_ends, &values).expect("Failed to create RunArray");
1537 let b = BooleanArray::from(vec![
1538 false, false, false, false, false, false, false, false, false, false,
1539 ]);
1540 let c = filter(&a, &b).unwrap();
1541 let actual: &RunArray<Int64Type> = as_run_array(&c);
1542 assert_eq!(0, actual.len());
1543 }
1544
1545 #[test]
1546 fn test_filter_run_end_encoding_array_max_value_gt_predicate_len() {
1547 let run_ends = Int64Array::from(vec![2, 3, 8, 10]);
1548 let values = Int64Array::from(vec![7, -2, 9, -8]);
1549 let a = RunArray::try_new(&run_ends, &values).expect("Failed to create RunArray");
1550 let b = BooleanArray::from(vec![false, true, true]);
1551 let c = filter(&a, &b).unwrap();
1552 let actual: &RunArray<Int64Type> = as_run_array(&c);
1553 assert_eq!(2, actual.len());
1554
1555 let expected = RunArray::try_new(
1556 &Int64Array::from(vec![1, 2]),
1557 &Int64Array::from(vec![7, -2]),
1558 )
1559 .expect("Failed to make expected RunArray test is broken");
1560
1561 assert_eq!(&actual.run_ends().values(), &expected.run_ends().values());
1562 assert_eq!(actual.values(), expected.values())
1563 }
1564
1565 #[test]
1566 fn test_filter_dictionary_array() {
1567 let values = [Some("hello"), None, Some("world"), Some("!")];
1568 let a: Int8DictionaryArray = values.iter().copied().collect();
1569 let b = BooleanArray::from(vec![false, true, true, false]);
1570 let c = filter(&a, &b).unwrap();
1571 let d = c
1572 .as_ref()
1573 .as_any()
1574 .downcast_ref::<Int8DictionaryArray>()
1575 .unwrap();
1576 let value_array = d.values();
1577 let values = value_array.as_any().downcast_ref::<StringArray>().unwrap();
1578 assert_eq!(3, values.len());
1580 assert_eq!(2, d.len());
1582 assert!(d.is_null(0));
1583 assert_eq!("world", values.value(d.keys().value(1) as usize));
1584 }
1585
1586 #[test]
1587 fn test_filter_list_array() {
1588 let field = Arc::new(Field::new_list_field(DataType::Int32, false));
1589 let offsets = OffsetBuffer::new(vec![0i64, 3, 6, 8, 8].into());
1590 let value_array = Arc::new(Int32Array::from_iter_values(0..8));
1591 let nulls = Some(NullBuffer::from(vec![true, true, true, false]));
1592 let a = LargeListArray::new(field.clone(), offsets, value_array, nulls);
1594 let b = BooleanArray::from(vec![false, true, false, true]);
1595 let result = filter(&a, &b).unwrap();
1596
1597 let offsets = OffsetBuffer::new(vec![0i64, 3, 3].into());
1599 let value_array = Arc::new(Int32Array::from_iter_values([3, 4, 5]));
1600 let nulls = Some(NullBuffer::from(vec![true, false]));
1601 let expected: ArrayRef = Arc::new(LargeListArray::new(field, offsets, value_array, nulls));
1602
1603 assert_eq!(&expected, &result);
1604 }
1605
1606 fn test_case_filter_list_view<T: OffsetSizeTrait>() {
1607 let mut list_array = GenericListViewBuilder::<T, _>::new(Int32Builder::new());
1609 list_array.append_value([Some(1), Some(2)]);
1610 list_array.append_null();
1611 list_array.append_value([]);
1612 list_array.append_value([Some(3), Some(4)]);
1613
1614 let list_array = list_array.finish();
1615 let predicate = BooleanArray::from_iter([true, false, true, false]);
1616
1617 let filtered = filter(&list_array, &predicate)
1619 .unwrap()
1620 .as_list_view::<T>()
1621 .clone();
1622
1623 let mut expected =
1624 GenericListViewBuilder::<T, _>::with_capacity(Int32Builder::with_capacity(5), 3);
1625 expected.append_value([Some(1), Some(2)]);
1626 expected.append_value([]);
1627 let expected = expected.finish();
1628
1629 assert_eq!(&filtered, &expected);
1630 }
1631
1632 fn test_case_filter_sliced_list_view<T: OffsetSizeTrait>() {
1633 let mut list_array =
1635 GenericListViewBuilder::<T, _>::with_capacity(Int32Builder::with_capacity(6), 4);
1636 list_array.append_value([Some(1), Some(2)]);
1637 list_array.append_null();
1638 list_array.append_value([]);
1639 list_array.append_value([Some(3), Some(4)]);
1640
1641 let list_array = list_array.finish();
1642
1643 let sliced = list_array.slice(1, 3);
1645 let predicate = BooleanArray::from_iter([false, false, true]);
1646
1647 let filtered = filter(&sliced, &predicate)
1649 .unwrap()
1650 .as_list_view::<T>()
1651 .clone();
1652
1653 let mut expected = GenericListViewBuilder::<T, _>::new(Int32Builder::new());
1654 expected.append_value([Some(3), Some(4)]);
1655 let expected = expected.finish();
1656
1657 assert_eq!(&filtered, &expected);
1658 }
1659
1660 #[test]
1661 fn test_filter_list_view_array() {
1662 test_case_filter_list_view::<i32>();
1663 test_case_filter_list_view::<i64>();
1664
1665 test_case_filter_sliced_list_view::<i32>();
1666 test_case_filter_sliced_list_view::<i64>();
1667 }
1668
1669 #[test]
1670 fn test_slice_iterator_bits() {
1671 let filter_values = (0..64).map(|i| i == 1).collect::<Vec<bool>>();
1672 let filter = BooleanArray::from(filter_values);
1673 let filter_count = filter.true_count();
1674
1675 let iter = SlicesIterator::new(&filter);
1676 let chunks = iter.collect::<Vec<_>>();
1677
1678 assert_eq!(chunks, vec![(1, 2)]);
1679 assert_eq!(filter_count, 1);
1680 }
1681
1682 #[test]
1683 fn test_slice_iterator_bits1() {
1684 let filter_values = (0..64).map(|i| i != 1).collect::<Vec<bool>>();
1685 let filter = BooleanArray::from(filter_values);
1686 let filter_count = filter.true_count();
1687
1688 let iter = SlicesIterator::new(&filter);
1689 let chunks = iter.collect::<Vec<_>>();
1690
1691 assert_eq!(chunks, vec![(0, 1), (2, 64)]);
1692 assert_eq!(filter_count, 64 - 1);
1693 }
1694
1695 #[test]
1696 fn test_slice_iterator_chunk_and_bits() {
1697 let filter_values = (0..130).map(|i| i % 62 != 0).collect::<Vec<bool>>();
1698 let filter = BooleanArray::from(filter_values);
1699 let filter_count = filter.true_count();
1700
1701 let iter = SlicesIterator::new(&filter);
1702 let chunks = iter.collect::<Vec<_>>();
1703
1704 assert_eq!(chunks, vec![(1, 62), (63, 124), (125, 130)]);
1705 assert_eq!(filter_count, 61 + 61 + 5);
1706 }
1707
1708 #[test]
1709 fn test_filter_selection_iterators() {
1710 let slices = [(0, 2), (4, 5)];
1711 let mut ranges = Vec::new();
1712 let selection: FilterSlices<'_> = FilterIterator::Materialized(slices.iter().copied());
1713 selection.for_each(|range| ranges.push(range));
1714 assert_eq!(ranges, slices);
1715
1716 let filter = BooleanArray::from(vec![true, true, false, false, true]);
1717 let mut ranges = Vec::new();
1718 let selection: FilterSlices<'_> = FilterIterator::Lazy(SlicesIterator::new(&filter));
1719 selection
1720 .try_for_each(|range| {
1721 ranges.push(range);
1722 Ok::<(), ArrowError>(())
1723 })
1724 .unwrap();
1725 assert_eq!(ranges, vec![(0, 2), (4, 5)]);
1726
1727 let indices = [1, 3, 5];
1728 let mut selected = Vec::new();
1729 let selection: FilterIndices<'_> = FilterIterator::Materialized(indices.iter().copied());
1730 selection.for_each(|idx| selected.push(idx));
1731 assert_eq!(selected, indices);
1732
1733 let filter = BooleanArray::from(vec![false, true, false, true]);
1734 let mut selected = Vec::new();
1735 let selection: FilterIndices<'_> = FilterIterator::Lazy(IndexIterator::new(&filter, 2));
1736 selection
1737 .try_for_each(|idx| {
1738 selected.push(idx);
1739 Ok::<(), ArrowError>(())
1740 })
1741 .unwrap();
1742 assert_eq!(selected, vec![1, 3]);
1743 }
1744
1745 #[test]
1746 fn test_null_mask() {
1747 let a = Int64Array::from(vec![Some(1), Some(2), None]);
1748
1749 let mask1 = BooleanArray::from(vec![Some(true), Some(true), None]);
1750 let out = filter(&a, &mask1).unwrap();
1751 assert_eq!(out.as_ref(), &a.slice(0, 2));
1752 }
1753
1754 #[test]
1755 fn test_filter_record_batch_no_columns() {
1756 let pred = BooleanArray::from(vec![Some(true), Some(true), None]);
1757 let options = RecordBatchOptions::default().with_row_count(Some(100));
1758 let record_batch =
1759 RecordBatch::try_new_with_options(Arc::new(Schema::empty()), vec![], &options).unwrap();
1760 let out = filter_record_batch(&record_batch, &pred).unwrap();
1761
1762 assert_eq!(out.num_rows(), 2);
1763 }
1764
1765 #[test]
1766 fn test_fast_path() {
1767 let a: PrimitiveArray<Int64Type> = PrimitiveArray::from(vec![Some(1), Some(2), None]);
1768
1769 let mask = BooleanArray::from(vec![true, true, true]);
1771 let out = filter(&a, &mask).unwrap();
1772 let b = out
1773 .as_any()
1774 .downcast_ref::<PrimitiveArray<Int64Type>>()
1775 .unwrap();
1776 assert_eq!(&a, b);
1777
1778 let mask = BooleanArray::from(vec![false, false, false]);
1780 let out = filter(&a, &mask).unwrap();
1781 assert_eq!(out.len(), 0);
1782 assert_eq!(out.data_type(), &DataType::Int64);
1783 }
1784
1785 #[test]
1786 fn test_slices() {
1787 let bools = std::iter::repeat_n(true, 10)
1789 .chain(std::iter::repeat_n(false, 30))
1790 .chain(std::iter::repeat_n(true, 20))
1791 .chain(std::iter::repeat_n(false, 17))
1792 .chain(std::iter::repeat_n(true, 4));
1793
1794 let bool_array: BooleanArray = bools.map(Some).collect();
1795
1796 let slices: Vec<_> = SlicesIterator::new(&bool_array).collect();
1797 let expected = vec![(0, 10), (40, 60), (77, 81)];
1798 assert_eq!(slices, expected);
1799
1800 let len = bool_array.len();
1802 let sliced_array = bool_array.slice(7, len - 10);
1803 let sliced_array = sliced_array
1804 .as_any()
1805 .downcast_ref::<BooleanArray>()
1806 .unwrap();
1807 let slices: Vec<_> = SlicesIterator::new(sliced_array).collect();
1808 let expected = vec![(0, 3), (33, 53), (70, 71)];
1809 assert_eq!(slices, expected);
1810 }
1811
1812 fn test_slices_fuzz(mask_len: usize, offset: usize, truncate: usize) {
1813 let mut rng = rng();
1814
1815 let bools: Vec<bool> = std::iter::from_fn(|| Some(rng.random()))
1816 .take(mask_len)
1817 .collect();
1818
1819 let buffer = Buffer::from_iter(bools.iter().copied());
1820
1821 let truncated_length = mask_len - offset - truncate;
1822
1823 let filter = BooleanArray::new(BooleanBuffer::new(buffer, offset, truncated_length), None);
1824
1825 let slice_bits: Vec<_> = SlicesIterator::new(&filter)
1826 .flat_map(|(start, end)| start..end)
1827 .collect();
1828
1829 let count = filter.true_count();
1830 let index_bits: Vec<_> = IndexIterator::new(&filter, count).collect();
1831
1832 let expected_bits: Vec<_> = bools
1833 .iter()
1834 .skip(offset)
1835 .take(truncated_length)
1836 .enumerate()
1837 .filter_map(|(idx, v)| v.then_some(idx))
1838 .collect();
1839
1840 assert_eq!(slice_bits, expected_bits);
1841 assert_eq!(index_bits, expected_bits);
1842 }
1843
1844 #[test]
1845 #[cfg_attr(miri, ignore)] fn fuzz_test_slices_iterator() {
1847 let mut rng = rng();
1848
1849 let uusize = UniformUsize::new(usize::MIN, usize::MAX).unwrap();
1850 for _ in 0..100 {
1851 let mask_len = rng.random_range(0..1024);
1852 let max_offset = 64.min(mask_len);
1853 let offset = uusize.sample(&mut rng).checked_rem(max_offset).unwrap_or(0);
1854
1855 let max_truncate = 128.min(mask_len - offset);
1856 let truncate = uusize
1857 .sample(&mut rng)
1858 .checked_rem(max_truncate)
1859 .unwrap_or(0);
1860
1861 test_slices_fuzz(mask_len, offset, truncate);
1862 }
1863
1864 test_slices_fuzz(64, 0, 0);
1865 test_slices_fuzz(64, 8, 0);
1866 test_slices_fuzz(64, 8, 8);
1867 test_slices_fuzz(32, 8, 8);
1868 test_slices_fuzz(32, 5, 9);
1869 }
1870
1871 fn filter_rust<T>(values: impl IntoIterator<Item = T>, predicate: &[bool]) -> Vec<T> {
1873 values
1874 .into_iter()
1875 .zip(predicate)
1876 .filter(|(_, x)| **x)
1877 .map(|(a, _)| a)
1878 .collect()
1879 }
1880
1881 fn gen_primitive<T>(len: usize, valid_percent: f64) -> Vec<Option<T>>
1883 where
1884 StandardUniform: Distribution<T>,
1885 {
1886 let mut rng = rng();
1887 (0..len)
1888 .map(|_| rng.random_bool(valid_percent).then(|| rng.random()))
1889 .collect()
1890 }
1891
1892 fn gen_strings(
1894 len: usize,
1895 valid_percent: f64,
1896 str_len_range: std::ops::Range<usize>,
1897 ) -> Vec<Option<String>> {
1898 let mut rng = rng();
1899 (0..len)
1900 .map(|_| {
1901 rng.random_bool(valid_percent).then(|| {
1902 let len = rng.random_range(str_len_range.clone());
1903 (0..len)
1904 .map(|_| char::from(rng.sample(Alphanumeric)))
1905 .collect()
1906 })
1907 })
1908 .collect()
1909 }
1910
1911 fn as_deref<T: std::ops::Deref>(src: &[Option<T>]) -> impl Iterator<Item = Option<&T::Target>> {
1913 src.iter().map(|x| x.as_deref())
1914 }
1915
1916 #[test]
1917 #[cfg_attr(miri, ignore)] fn fuzz_filter() {
1919 let mut rng = rng();
1920
1921 for i in 0..100 {
1922 let filter_percent = match i {
1923 0..=4 => 1.,
1924 5..=10 => 0.,
1925 _ => rng.random_range(0.0..1.0),
1926 };
1927
1928 let valid_percent = rng.random_range(0.0..1.0);
1929
1930 let array_len = rng.random_range(32..256);
1931 let array_offset = rng.random_range(0..10);
1932
1933 let filter_offset = rng.random_range(0..10);
1935 let filter_truncate = rng.random_range(0..10);
1936 let bools: Vec<_> = std::iter::from_fn(|| Some(rng.random_bool(filter_percent)))
1937 .take(array_len + filter_offset - filter_truncate)
1938 .collect();
1939
1940 let predicate = BooleanArray::from_iter(bools.iter().copied().map(Some));
1941
1942 let predicate = predicate.slice(filter_offset, array_len - filter_truncate);
1944 let predicate = predicate.as_any().downcast_ref::<BooleanArray>().unwrap();
1945 let bools = &bools[filter_offset..];
1946
1947 let values = gen_primitive(array_len + array_offset, valid_percent);
1949 let src = Int32Array::from_iter(values.iter().copied());
1950
1951 let src = src.slice(array_offset, array_len);
1952 let src = src.as_any().downcast_ref::<Int32Array>().unwrap();
1953 let values = &values[array_offset..];
1954
1955 let filtered = filter(src, predicate).unwrap();
1956 let array = filtered.as_any().downcast_ref::<Int32Array>().unwrap();
1957 let actual: Vec<_> = array.iter().collect();
1958
1959 assert_eq!(actual, filter_rust(values.iter().copied(), bools));
1960
1961 let strings = gen_strings(array_len + array_offset, valid_percent, 0..20);
1963 let src = StringArray::from_iter(as_deref(&strings));
1964
1965 let src = src.slice(array_offset, array_len);
1966 let src = src.as_any().downcast_ref::<StringArray>().unwrap();
1967
1968 let filtered = filter(src, predicate).unwrap();
1969 let array = filtered.as_any().downcast_ref::<StringArray>().unwrap();
1970 let actual: Vec<_> = array.iter().collect();
1971
1972 let expected_strings = filter_rust(as_deref(&strings[array_offset..]), bools);
1973 assert_eq!(actual, expected_strings);
1974
1975 let src = DictionaryArray::<Int32Type>::from_iter(as_deref(&strings));
1977
1978 let src = src.slice(array_offset, array_len);
1979 let src = src
1980 .as_any()
1981 .downcast_ref::<DictionaryArray<Int32Type>>()
1982 .unwrap();
1983
1984 let filtered = filter(src, predicate).unwrap();
1985
1986 let array = filtered
1987 .as_any()
1988 .downcast_ref::<DictionaryArray<Int32Type>>()
1989 .unwrap();
1990
1991 let values = array
1992 .values()
1993 .as_any()
1994 .downcast_ref::<StringArray>()
1995 .unwrap();
1996
1997 let actual: Vec<_> = array
1998 .keys()
1999 .iter()
2000 .map(|key| key.map(|key| values.value(key as usize)))
2001 .collect();
2002
2003 assert_eq!(actual, expected_strings);
2004 }
2005 }
2006
2007 #[test]
2008 fn test_filter_map() {
2009 let mut builder =
2010 MapBuilder::new(None, StringBuilder::new(), Int64Builder::with_capacity(4));
2011 builder.keys().append_value("key1");
2013 builder.values().append_value(1);
2014 builder.append(true).unwrap();
2015 builder.keys().append_value("key2");
2016 builder.keys().append_value("key3");
2017 builder.values().append_value(2);
2018 builder.values().append_value(3);
2019 builder.append(true).unwrap();
2020 builder.append(false).unwrap();
2021 builder.keys().append_value("key1");
2022 builder.values().append_value(1);
2023 builder.append(true).unwrap();
2024 let maparray = Arc::new(builder.finish()) as ArrayRef;
2025
2026 let indices = vec![Some(true), Some(false), Some(false), Some(true)]
2027 .into_iter()
2028 .collect::<BooleanArray>();
2029 let got = filter(&maparray, &indices).unwrap();
2030
2031 let mut builder =
2032 MapBuilder::new(None, StringBuilder::new(), Int64Builder::with_capacity(2));
2033 builder.keys().append_value("key1");
2034 builder.values().append_value(1);
2035 builder.append(true).unwrap();
2036 builder.keys().append_value("key1");
2037 builder.values().append_value(1);
2038 builder.append(true).unwrap();
2039 let expected = Arc::new(builder.finish()) as ArrayRef;
2040
2041 assert_eq!(&expected, &got);
2042 }
2043
2044 #[test]
2045 fn test_filter_fixed_size_list_arrays() {
2046 let field = Arc::new(Field::new_list_field(DataType::Int32, false));
2047 let value_array = Arc::new(Int32Array::from_iter_values(0..9));
2048 let array = FixedSizeListArray::new(field, 3, value_array, None);
2049
2050 let filter_array = BooleanArray::from(vec![true, false, false]);
2051
2052 let c = filter(&array, &filter_array).unwrap();
2053 let filtered = c.as_any().downcast_ref::<FixedSizeListArray>().unwrap();
2054
2055 assert_eq!(filtered.len(), 1);
2056
2057 let list = filtered.value(0);
2058 assert_eq!(
2059 &[0, 1, 2],
2060 list.as_any().downcast_ref::<Int32Array>().unwrap().values()
2061 );
2062
2063 let filter_array = BooleanArray::from(vec![true, false, true]);
2064
2065 let c = filter(&array, &filter_array).unwrap();
2066 let filtered = c.as_any().downcast_ref::<FixedSizeListArray>().unwrap();
2067
2068 assert_eq!(filtered.len(), 2);
2069
2070 let list = filtered.value(0);
2071 assert_eq!(
2072 &[0, 1, 2],
2073 list.as_any().downcast_ref::<Int32Array>().unwrap().values()
2074 );
2075 let list = filtered.value(1);
2076 assert_eq!(
2077 &[6, 7, 8],
2078 list.as_any().downcast_ref::<Int32Array>().unwrap().values()
2079 );
2080 }
2081
2082 #[test]
2083 fn test_filter_fixed_size_list_arrays_with_null() {
2084 let field = Arc::new(Field::new_list_field(DataType::Int32, false));
2085 let value_array = Arc::new(Int32Array::from_iter_values(0..10));
2086 let nulls = Some(NullBuffer::from(vec![true, false, false, true, true]));
2087 let array = FixedSizeListArray::new(field, 2, value_array, nulls);
2088
2089 let filter_array = BooleanArray::from(vec![true, true, false, true, false]);
2090
2091 let c = filter(&array, &filter_array).unwrap();
2092 let filtered = c.as_any().downcast_ref::<FixedSizeListArray>().unwrap();
2093
2094 assert_eq!(filtered.len(), 3);
2095
2096 let list = filtered.value(0);
2097 assert_eq!(
2098 &[0, 1],
2099 list.as_any().downcast_ref::<Int32Array>().unwrap().values()
2100 );
2101 assert!(filtered.is_null(1));
2102 let list = filtered.value(2);
2103 assert_eq!(
2104 &[6, 7],
2105 list.as_any().downcast_ref::<Int32Array>().unwrap().values()
2106 );
2107 }
2108
2109 fn test_filter_union_array(array: UnionArray) {
2110 let filter_array = BooleanArray::from(vec![true, false, false]);
2111 let c = filter(&array, &filter_array).unwrap();
2112 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2113
2114 let mut builder = UnionBuilder::new_dense();
2115 builder.append::<Int32Type>("A", 1).unwrap();
2116 let expected_array = builder.build().unwrap();
2117
2118 compare_union_arrays(filtered, &expected_array);
2119
2120 let filter_array = BooleanArray::from(vec![true, false, true]);
2121 let c = filter(&array, &filter_array).unwrap();
2122 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2123
2124 let mut builder = UnionBuilder::new_dense();
2125 builder.append::<Int32Type>("A", 1).unwrap();
2126 builder.append::<Int32Type>("A", 34).unwrap();
2127 let expected_array = builder.build().unwrap();
2128
2129 compare_union_arrays(filtered, &expected_array);
2130
2131 let filter_array = BooleanArray::from(vec![true, true, false]);
2132 let c = filter(&array, &filter_array).unwrap();
2133 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2134
2135 let mut builder = UnionBuilder::new_dense();
2136 builder.append::<Int32Type>("A", 1).unwrap();
2137 builder.append::<Float64Type>("B", 3.2).unwrap();
2138 let expected_array = builder.build().unwrap();
2139
2140 compare_union_arrays(filtered, &expected_array);
2141 }
2142
2143 #[test]
2144 fn test_filter_union_array_dense() {
2145 let mut builder = UnionBuilder::new_dense();
2146 builder.append::<Int32Type>("A", 1).unwrap();
2147 builder.append::<Float64Type>("B", 3.2).unwrap();
2148 builder.append::<Int32Type>("A", 34).unwrap();
2149 let array = builder.build().unwrap();
2150
2151 test_filter_union_array(array);
2152 }
2153
2154 #[test]
2155 fn test_filter_run_union_array_dense() {
2156 let mut builder = UnionBuilder::new_dense();
2157 builder.append::<Int32Type>("A", 1).unwrap();
2158 builder.append::<Int32Type>("A", 3).unwrap();
2159 builder.append::<Int32Type>("A", 34).unwrap();
2160 let array = builder.build().unwrap();
2161
2162 let filter_array = BooleanArray::from(vec![true, true, false]);
2163 let c = filter(&array, &filter_array).unwrap();
2164 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2165
2166 let mut builder = UnionBuilder::new_dense();
2167 builder.append::<Int32Type>("A", 1).unwrap();
2168 builder.append::<Int32Type>("A", 3).unwrap();
2169 let expected = builder.build().unwrap();
2170
2171 assert_eq!(filtered.to_data(), expected.to_data());
2172 }
2173
2174 #[test]
2175 fn test_filter_union_array_dense_with_nulls() {
2176 let mut builder = UnionBuilder::new_dense();
2177 builder.append::<Int32Type>("A", 1).unwrap();
2178 builder.append::<Float64Type>("B", 3.2).unwrap();
2179 builder.append_null::<Float64Type>("B").unwrap();
2180 builder.append::<Int32Type>("A", 34).unwrap();
2181 let array = builder.build().unwrap();
2182
2183 let filter_array = BooleanArray::from(vec![true, true, false, false]);
2184 let c = filter(&array, &filter_array).unwrap();
2185 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2186
2187 let mut builder = UnionBuilder::new_dense();
2188 builder.append::<Int32Type>("A", 1).unwrap();
2189 builder.append::<Float64Type>("B", 3.2).unwrap();
2190 let expected_array = builder.build().unwrap();
2191
2192 compare_union_arrays(filtered, &expected_array);
2193
2194 let filter_array = BooleanArray::from(vec![true, false, true, false]);
2195 let c = filter(&array, &filter_array).unwrap();
2196 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2197
2198 let mut builder = UnionBuilder::new_dense();
2199 builder.append::<Int32Type>("A", 1).unwrap();
2200 builder.append_null::<Float64Type>("B").unwrap();
2201 let expected_array = builder.build().unwrap();
2202
2203 compare_union_arrays(filtered, &expected_array);
2204 }
2205
2206 #[test]
2207 fn test_filter_union_array_sparse() {
2208 let mut builder = UnionBuilder::new_sparse();
2209 builder.append::<Int32Type>("A", 1).unwrap();
2210 builder.append::<Float64Type>("B", 3.2).unwrap();
2211 builder.append::<Int32Type>("A", 34).unwrap();
2212 let array = builder.build().unwrap();
2213
2214 test_filter_union_array(array);
2215 }
2216
2217 #[test]
2218 fn test_filter_union_array_sparse_with_nulls() {
2219 let mut builder = UnionBuilder::new_sparse();
2220 builder.append::<Int32Type>("A", 1).unwrap();
2221 builder.append::<Float64Type>("B", 3.2).unwrap();
2222 builder.append_null::<Float64Type>("B").unwrap();
2223 builder.append::<Int32Type>("A", 34).unwrap();
2224 let array = builder.build().unwrap();
2225
2226 let filter_array = BooleanArray::from(vec![true, false, true, false]);
2227 let c = filter(&array, &filter_array).unwrap();
2228 let filtered = c.as_any().downcast_ref::<UnionArray>().unwrap();
2229
2230 let mut builder = UnionBuilder::new_sparse();
2231 builder.append::<Int32Type>("A", 1).unwrap();
2232 builder.append_null::<Float64Type>("B").unwrap();
2233 let expected_array = builder.build().unwrap();
2234
2235 compare_union_arrays(filtered, &expected_array);
2236 }
2237
2238 fn compare_union_arrays(union1: &UnionArray, union2: &UnionArray) {
2239 assert_eq!(union1.len(), union2.len());
2240
2241 for i in 0..union1.len() {
2242 let type_id = union1.type_id(i);
2243
2244 let slot1 = union1.value(i);
2245 let slot2 = union2.value(i);
2246
2247 assert_eq!(slot1.is_null(0), slot2.is_null(0));
2248
2249 if !slot1.is_null(0) && !slot2.is_null(0) {
2250 match type_id {
2251 0 => {
2252 let slot1 = slot1.as_any().downcast_ref::<Int32Array>().unwrap();
2253 assert_eq!(slot1.len(), 1);
2254 let value1 = slot1.value(0);
2255
2256 let slot2 = slot2.as_any().downcast_ref::<Int32Array>().unwrap();
2257 assert_eq!(slot2.len(), 1);
2258 let value2 = slot2.value(0);
2259 assert_eq!(value1, value2);
2260 }
2261 1 => {
2262 let slot1 = slot1.as_any().downcast_ref::<Float64Array>().unwrap();
2263 assert_eq!(slot1.len(), 1);
2264 let value1 = slot1.value(0);
2265
2266 let slot2 = slot2.as_any().downcast_ref::<Float64Array>().unwrap();
2267 assert_eq!(slot2.len(), 1);
2268 let value2 = slot2.value(0);
2269 assert_eq!(value1, value2);
2270 }
2271 _ => unreachable!(),
2272 }
2273 }
2274 }
2275 }
2276
2277 #[test]
2278 fn test_filter_struct() {
2279 let predicate = BooleanArray::from(vec![true, false, true, false]);
2280
2281 let a = Arc::new(StringArray::from(vec!["hello", " ", "world", "!"]));
2282 let a_filtered = Arc::new(StringArray::from(vec!["hello", "world"]));
2283
2284 let b = Arc::new(Int32Array::from(vec![5, 6, 7, 8]));
2285 let b_filtered = Arc::new(Int32Array::from(vec![5, 7]));
2286
2287 let null_mask = NullBuffer::from(vec![true, false, false, true]);
2288 let null_mask_filtered = NullBuffer::from(vec![true, false]);
2289
2290 let a_field = Field::new("a", DataType::Utf8, false);
2291 let b_field = Field::new("b", DataType::Int32, false);
2292
2293 let array = StructArray::new(vec![a_field.clone()].into(), vec![a.clone()], None);
2294 let expected =
2295 StructArray::new(vec![a_field.clone()].into(), vec![a_filtered.clone()], None);
2296
2297 let result = filter(&array, &predicate).unwrap();
2298
2299 assert_eq!(result.to_data(), expected.to_data());
2300
2301 let array = StructArray::new(
2302 vec![a_field.clone()].into(),
2303 vec![a.clone()],
2304 Some(null_mask.clone()),
2305 );
2306 let expected = StructArray::new(
2307 vec![a_field.clone()].into(),
2308 vec![a_filtered.clone()],
2309 Some(null_mask_filtered.clone()),
2310 );
2311
2312 let result = filter(&array, &predicate).unwrap();
2313
2314 assert_eq!(result.to_data(), expected.to_data());
2315
2316 let array = StructArray::new(
2317 vec![a_field.clone(), b_field.clone()].into(),
2318 vec![a.clone(), b.clone()],
2319 None,
2320 );
2321 let expected = StructArray::new(
2322 vec![a_field.clone(), b_field.clone()].into(),
2323 vec![a_filtered.clone(), b_filtered.clone()],
2324 None,
2325 );
2326
2327 let result = filter(&array, &predicate).unwrap();
2328
2329 assert_eq!(result.to_data(), expected.to_data());
2330
2331 let array = StructArray::new(
2332 vec![a_field.clone(), b_field.clone()].into(),
2333 vec![a.clone(), b.clone()],
2334 Some(null_mask.clone()),
2335 );
2336
2337 let expected = StructArray::new(
2338 vec![a_field.clone(), b_field.clone()].into(),
2339 vec![a_filtered.clone(), b_filtered.clone()],
2340 Some(null_mask_filtered.clone()),
2341 );
2342
2343 let result = filter(&array, &predicate).unwrap();
2344
2345 assert_eq!(result.to_data(), expected.to_data());
2346 }
2347
2348 #[test]
2349 fn test_filter_empty_struct() {
2350 let fields = arrow_schema::Field::new(
2357 "a",
2358 arrow_schema::DataType::Struct(arrow_schema::Fields::from(vec![
2359 arrow_schema::Field::new("b", arrow_schema::DataType::Int64, true),
2360 arrow_schema::Field::new(
2361 "c",
2362 arrow_schema::DataType::Struct(arrow_schema::Fields::empty()),
2363 true,
2364 ),
2365 ])),
2366 true,
2367 );
2368
2369 let schema = Arc::new(Schema::new(vec![fields]));
2377
2378 let b = Arc::new(Int64Array::from(vec![None, None, None]));
2379 let c = Arc::new(StructArray::new_empty_fields(
2380 3,
2381 Some(NullBuffer::from(vec![true, true, true])),
2382 ));
2383 let a = StructArray::new(
2384 vec![
2385 Field::new("b", DataType::Int64, true),
2386 Field::new("c", DataType::Struct(Fields::empty()), true),
2387 ]
2388 .into(),
2389 vec![b.clone(), c.clone()],
2390 Some(NullBuffer::from(vec![true, true, true])),
2391 );
2392 let record_batch = RecordBatch::try_new(schema, vec![Arc::new(a)]).unwrap();
2393 println!("{record_batch:?}");
2394
2395 let predicate = BooleanArray::from(vec![true, false, true]);
2397 let filtered_batch = filter_record_batch(&record_batch, &predicate).unwrap();
2398
2399 assert_eq!(filtered_batch.num_rows(), 2);
2401 }
2402
2403 #[test]
2404 #[should_panic(expected = "buffer.len() >= predicate.filter.len()")]
2405 fn test_filter_bits_too_large() {
2406 let buffer = BooleanBuffer::from(vec![false; 8]);
2407 let predicate = BooleanArray::from(vec![true; 9]);
2408 let filter = FilterBuilder::new(&predicate).build();
2409 filter_bits(&buffer, &filter);
2410 }
2411
2412 #[test]
2413 #[should_panic(expected = "values.len() >= predicate.filter.len()")]
2414 fn test_filter_native_too_large() {
2415 let values = vec![1; 8];
2416 let predicate = BooleanArray::from(vec![false; 9]);
2417 let filter = FilterBuilder::new(&predicate).build();
2418 filter_native(&values, &filter);
2419 }
2420}