Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 20 additions & 2 deletions datafusion/physical-expr/src/expressions/in_list.rs
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ use datafusion_expr::{ColumnarValue, expr_vec_fmt};

mod array_static_filter;
mod branchless_filter;
mod byte_view_filter;
mod primitive_filter;
mod result;
mod static_filter;
Expand Down Expand Up @@ -215,7 +216,7 @@ impl InListExpr {
expr,
list,
negated,
Some(instantiate_static_filter(array)?),
Some(instantiate_static_filter(array, &expr_data_type)?),
))
}

Expand All @@ -242,7 +243,7 @@ impl InListExpr {

// Try to create a static filter if all list expressions are constants
let static_filter = match try_evaluate_constant_list(&list, schema)? {
Some(in_array) => Some(instantiate_static_filter(in_array)?),
Some(in_array) => Some(instantiate_static_filter(in_array, &expr_data_type)?),
None => None, // Non-constant expressions, fall back to dynamic evaluation
};

Expand Down Expand Up @@ -3576,6 +3577,23 @@ mod tests {
)?
);

// Utf8View in_array, Utf8View and Dict(Utf8View) needles
let utf8view_in =
Arc::new(StringViewArray::from(vec!["a", "b", "c"])) as ArrayRef;
let utf8view_needle =
Arc::new(StringViewArray::from(vec!["a", "d", "b"])) as ArrayRef;
assert_eq!(
expected,
eval_in_list_from_array(
Arc::clone(&utf8view_needle),
Arc::clone(&utf8view_in),
)?
);
assert_eq!(
expected,
eval_in_list_from_array(wrap_in_dict(utf8view_needle), utf8view_in)?
);

// Struct in_array, Struct needle: multi-column join
let struct_fields = Fields::from(vec![
Field::new("c0", DataType::Utf8, true),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@
use std::mem::size_of;

use arrow::array::{Array, ArrayRef, AsArray, BooleanArray, PrimitiveArray};
use arrow::buffer::{BooleanBuffer, ScalarBuffer};
use arrow::buffer::{BooleanBuffer, NullBuffer, ScalarBuffer};
use arrow::datatypes::*;
use arrow::util::bit_iterator::BitIndexIterator;
use datafusion_common::{Result, exec_datafusion_err, internal_datafusion_err};
Expand Down Expand Up @@ -244,6 +244,17 @@ where
check_values,
})
}

#[inline]
pub(super) fn contains_slice(
&self,
input_values: &[BranchlessNative<T>],
nulls: Option<&NullBuffer>,
negated: bool,
) -> BooleanArray {
let matches = (self.check_values)(self.in_list_values.as_ref(), input_values);
build_result_from_contains(nulls, self.null_count > 0, negated, matches)
}
}

impl<T> StaticFilter for BranchlessFilter<T>
Expand Down Expand Up @@ -272,14 +283,7 @@ where
exec_datafusion_err!("BranchlessFilter: expected {} array", T::DATA_TYPE)
})?;
let input_values = branchless_values::<T>(v);
let matches =
(self.check_values)(self.in_list_values.as_ref(), input_values.as_ref());
Ok(build_result_from_contains(
v.nulls(),
self.null_count > 0,
negated,
matches,
))
Ok(self.contains_slice(input_values.as_ref(), v.nulls(), negated))
}
}

Expand Down
Loading
Loading