From 14062da55fe9c3807c355e6270cae7e9ee40ace3 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 10:05:25 -0400 Subject: [PATCH 1/6] Add SQL null semantics to list_contains Add explicit SQL three-valued membership semantics, including nullable return types, expression helpers, option serialization, and propagation through sequence kernels. Use SQL semantics for DataFusion IN conversion. Retain the existing comparison/OR evaluation for constant sets. Set preparation and single-pass probes follow in the next change. Include null-semantics regressions for scalar and column inputs, empty and null lists, sequence fallback, serialization, and DataFusion. Signed-off-by: Robert Kruszewski --- .../sequence/src/compute/list_contains.rs | 75 +- vortex-array/proto/expr.proto | 4 + vortex-array/src/builtins.rs | 3 +- vortex-array/src/expr/exprs.rs | 55 +- vortex-array/src/expr/mod.rs | 2 + .../src/scalar_fn/fns/list_contains/kernel.rs | 7 +- .../src/scalar_fn/fns/list_contains/mod.rs | 663 +++++++++++++++--- vortex-datafusion/src/convert/exprs.rs | 98 ++- 8 files changed, 801 insertions(+), 106 deletions(-) diff --git a/encodings/sequence/src/compute/list_contains.rs b/encodings/sequence/src/compute/list_contains.rs index 80ffcad24cd..5b3e620a71c 100644 --- a/encodings/sequence/src/compute/list_contains.rs +++ b/encodings/sequence/src/compute/list_contains.rs @@ -8,7 +8,7 @@ use vortex_array::arrays::BoolArray; use vortex_array::arrays::ConstantArray; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce; -use vortex_error::VortexExpect; +use vortex_array::scalar_fn::fns::list_contains::ListContainsOptions; use vortex_error::VortexResult; use crate::array::Sequence; @@ -19,17 +19,27 @@ impl ListContainsElementReduce for Sequence { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, + options: &ListContainsOptions, ) -> VortexResult> { let Some(list_scalar) = list.as_constant() else { return Ok(None); }; - let list_elements = list_scalar - .as_list() - .elements() - .vortex_expect("non-null element (checked in entry)"); + // A null list scalar has no elements to intersect with. Nothing checks this before the + // reduce rule runs, so fall back to the generic implementation, which resolves a null + // haystack to all-null rather than panicking here. + let Some(list_elements) = list_scalar.as_list().elements() else { + return Ok(None); + }; + + // The intersection search treats a null element as matching nothing, which under SQL null + // semantics is only half the answer: there a non-match is unknown, which the search cannot + // express. + if options.sql_null_semantics && list_elements.iter().any(Scalar::is_null) { + return Ok(None); + } - let nullability = list.dtype().nullability() | element.dtype().nullability(); + let nullability = options.result_nullability(list.dtype(), element.dtype()); let mut set_indices: Vec = Vec::new(); for intercept in list_elements.iter() { @@ -69,12 +79,15 @@ mod tests { use vortex_array::VortexSessionExecute; use vortex_array::arrays::BoolArray; use vortex_array::assert_arrays_eq; + use vortex_array::dtype::DType; use vortex_array::dtype::Nullability; use vortex_array::dtype::PType::I32; use vortex_array::expr::list_contains; + use vortex_array::expr::list_contains_opts; use vortex_array::expr::lit; use vortex_array::expr::root; use vortex_array::scalar::Scalar; + use vortex_array::scalar_fn::fns::list_contains::ListContainsOptions; use vortex_session::VortexSession; use crate::Sequence; @@ -122,6 +135,56 @@ mod tests { } } + #[test] + fn test_list_contains_null_list() { + // A null haystack resolves to null for every row. The reduce rule used to assume the list + // scalar was non-null and panicked instead of declining the reduction. + let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3) + .unwrap() + .into_array(); + + let null_list = Scalar::null(DType::List(Arc::new(I32.into()), Nullability::Nullable)); + let expr = list_contains(lit(null_list), root()); + let result = array.apply(&expr).unwrap(); + let expected = BoolArray::from_iter([None::, None, None]); + assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx()); + } + + #[test] + fn test_list_contains_null_element_semantics() { + // The sequence kernel skips a null element, which is right by default. Under SQL null + // semantics a non-match must be null instead, so a constant list holding a null goes to + // the generic path instead. + let element = DType::Primitive(I32, Nullability::Nullable); + let set = Scalar::list( + Arc::new(element.clone()), + vec![ + Scalar::primitive(1i32, Nullability::Nullable), + Scalar::null(element), + ], + Nullability::NonNullable, + ); + let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3) + .unwrap() + .into_array(); + + let result = array + .clone() + .apply(&list_contains(lit(set.clone()), root())) + .unwrap(); + let expected = BoolArray::from_iter([true, false, false]); + assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx()); + + let sql = ListContainsOptions { + sql_null_semantics: true, + }; + let result = array + .apply(&list_contains_opts(lit(set), root(), sql)) + .unwrap(); + let expected = BoolArray::from_iter([Some(true), None, None]); + assert_arrays_eq!(result, expected, &mut SESSION.create_execution_ctx()); + } + #[test] fn test_list_contains_constant_sequence() { let list_scalar = Scalar::list( diff --git a/vortex-array/proto/expr.proto b/vortex-array/proto/expr.proto index 00ffaac433c..a53a75ad50c 100644 --- a/vortex-array/proto/expr.proto +++ b/vortex-array/proto/expr.proto @@ -102,6 +102,10 @@ message LikeOpts { bool case_insensitive = 2; } +message ListContainsOpts { + bool sql_null_semantics = 1; +} + message CastOpts { vortex.dtype.DType target = 1; } diff --git a/vortex-array/src/builtins.rs b/vortex-array/src/builtins.rs index eb26931e753..42b696c2f47 100644 --- a/vortex-array/src/builtins.rs +++ b/vortex-array/src/builtins.rs @@ -31,6 +31,7 @@ use crate::scalar_fn::fns::get_item::GetItem; use crate::scalar_fn::fns::is_not_null::IsNotNull; use crate::scalar_fn::fns::is_null::IsNull; use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; use crate::scalar_fn::fns::mask::Mask; use crate::scalar_fn::fns::not::Not; use crate::scalar_fn::fns::operators::Operator; @@ -103,7 +104,7 @@ impl ExprBuiltins for Expression { } fn list_contains(&self, value: Expression) -> VortexResult { - ListContains.try_new_expr(EmptyOptions, [self.clone(), value]) + ListContains.try_new_expr(ListContainsOptions::default(), [self.clone(), value]) } fn zip(&self, if_true: Expression, if_false: Expression) -> VortexResult { diff --git a/vortex-array/src/expr/exprs.rs b/vortex-array/src/expr/exprs.rs index f860c12c040..9f414d39ea6 100644 --- a/vortex-array/src/expr/exprs.rs +++ b/vortex-array/src/expr/exprs.rs @@ -40,6 +40,7 @@ use crate::scalar_fn::fns::is_null::IsNull; use crate::scalar_fn::fns::like::Like; use crate::scalar_fn::fns::like::LikeOptions; use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; use crate::scalar_fn::fns::list_length::ListLength; use crate::scalar_fn::fns::list_sum::ListSum; use crate::scalar_fn::fns::literal::Literal; @@ -1120,23 +1121,64 @@ pub fn bound_dynamic( // ---- ListContains ---- -/// Creates an expression that checks if a value is contained in a list. +/// Creates an expression that checks whether a list contains a value. /// -/// Returns a boolean array indicating whether the value appears in each list. +/// A null element never matches anything: a needle that matches no element yields `false`. For +/// SQL's three-valued `IN`, where a null element makes that `null`, use [`in_list`]. /// /// ```rust /// # use vortex_array::expr::{list_contains, lit, root}; /// let expr = list_contains(root(), lit(42)); /// ``` pub fn list_contains(list: Expression, value: Expression) -> Expression { - ListContains.new_expr(EmptyOptions, [list, value]) + list_contains_opts(list, value, ListContainsOptions::default()) } -/// Creates a bound expression that checks if a value is contained in a list. +/// Creates an expression that checks whether a list contains a value, with explicit +/// [`ListContainsOptions`]. +pub fn list_contains_opts( + list: Expression, + value: Expression, + options: ListContainsOptions, +) -> Expression { + ListContains.new_expr(options, [list, value]) +} + +/// Creates `needle IN (list)` with SQL null semantics: a null element is an unknown value, so a +/// needle that matches no element yields `null` rather than `false` when the list holds one, and +/// `not(in_list(..))` never admits such a row. +/// +/// ```rust +/// # use vortex_array::expr::{in_list, lit, root}; +/// let expr = in_list(root(), lit(vec![1, 2, 3])); +/// ``` +pub fn in_list(needle: Expression, list: Expression) -> Expression { + list_contains_opts( + list, + needle, + ListContainsOptions { + sql_null_semantics: true, + }, + ) +} + +/// Creates a bound expression that checks whether a list contains a value. pub fn bound_list_contains(list: BoundExpression, value: BoundExpression) -> BoundExpression { ListContains - .try_new_bound_expr(EmptyOptions, [list, value]) - .vortex_expect("list-contains expressions require a compatible list and value dtype") + .try_new_bound_expr(ListContainsOptions::default(), [list, value]) + .vortex_expect("list-contains expressions require a list child") +} + +/// Creates a bound `needle IN (list)` with SQL null semantics. See [`in_list`]. +pub fn bound_in_list(needle: BoundExpression, list: BoundExpression) -> BoundExpression { + ListContains + .try_new_bound_expr( + ListContainsOptions { + sql_null_semantics: true, + }, + [list, needle], + ) + .vortex_expect("list-contains expressions require a list child") } // ---- ByteLength ---- @@ -1265,6 +1307,7 @@ pub mod bound { pub use super::bound_gt as gt; pub use super::bound_gt_eq as gt_eq; pub use super::bound_ilike as ilike; + pub use super::bound_in_list as in_list; pub use super::bound_is_nan as is_nan; pub use super::bound_is_not_null as is_not_null; pub use super::bound_is_null as is_null; diff --git a/vortex-array/src/expr/mod.rs b/vortex-array/src/expr/mod.rs index b5634728589..3112d7d79fb 100644 --- a/vortex-array/src/expr/mod.rs +++ b/vortex-array/src/expr/mod.rs @@ -79,12 +79,14 @@ pub use exprs::get_item; pub use exprs::gt; pub use exprs::gt_eq; pub use exprs::ilike; +pub use exprs::in_list; pub use exprs::is_nan; pub use exprs::is_not_null; pub use exprs::is_null; pub use exprs::is_root; pub use exprs::like; pub use exprs::list_contains; +pub use exprs::list_contains_opts; pub use exprs::list_length; pub use exprs::list_sum; pub use exprs::list_sum_opts; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs index 563600bfeee..7cee5f7aac1 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -15,6 +15,7 @@ use crate::arrays::scalar_fn::ScalarFnArrayView; use crate::kernel::ExecuteParentKernel; use crate::optimizer::rules::ArrayParentReduceRule; use crate::scalar_fn::fns::list_contains::ListContains as ListContainsExpr; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; /// Check list-contains without reading buffers (metadata-only). /// @@ -30,6 +31,7 @@ pub trait ListContainsElementReduce: VTable { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, + options: &ListContainsOptions, ) -> VortexResult>; } @@ -42,6 +44,7 @@ pub trait ListContainsElementKernel: VTable { fn list_contains( list: &ArrayRef, element: ArrayView<'_, Self>, + options: &ListContainsOptions, ctx: &mut ExecutionCtx, ) -> VortexResult>; } @@ -70,7 +73,7 @@ where .as_opt::() .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); let list = scalar_fn_array.get_child(0); - ::list_contains(list, array) + ::list_contains(list, array, parent.options) } } @@ -99,6 +102,6 @@ where .as_opt::() .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); let list = scalar_fn_array.get_child(0); - ::list_contains(list, array, ctx) + ::list_contains(list, array, parent.options, ctx) } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index 2a5e8ea19e6..d3182d2fdb3 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -3,11 +3,14 @@ mod kernel; +use std::fmt::Display; +use std::fmt::Formatter; use std::ops::BitOr; use arrow_buffer::bit_iterator::BitIndexIterator; pub use kernel::*; use num_traits::Zero; +use prost::Message; use vortex_buffer::BitBuffer; use vortex_error::VortexExpect; use vortex_error::VortexResult; @@ -35,11 +38,11 @@ use crate::dtype::IntegerPType; use crate::dtype::Nullability; use crate::match_each_integer_ptype; use crate::match_each_unsigned_integer_ptype; +use crate::proto::expr as pb; use crate::scalar::ListScalar; use crate::scalar::Scalar; use crate::scalar_fn::Arity; use crate::scalar_fn::ChildName; -use crate::scalar_fn::EmptyOptions; use crate::scalar_fn::ExecutionArgs; use crate::scalar_fn::ScalarFnId; use crate::scalar_fn::ScalarFnVTable; @@ -51,35 +54,102 @@ use crate::validity::Validity; #[derive(Clone)] pub struct ListContains; +/// Options for [`ListContains`]. +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)] +pub struct ListContainsOptions { + /// Whether a null list element is SQL's unknown value. + /// + /// Off, the default, a null element is never equal to anything: a needle that matches no + /// element yields `false`, and a needle that matches some element yields `true`. On, the + /// comparison against a null element is unknown, so a needle that matches no element yields + /// `null` when the list holds a null — the three-valued semantics of SQL `IN`, under which + /// `x NOT IN (1, NULL)` is never true. + /// + /// A null needle yields `null` against a list with elements either way. Against an empty list + /// there is nothing to compare it to: off, that is `false`; on, a null needle is `null` against + /// any list, which makes [`ListContains`] strict. + pub sql_null_semantics: bool, +} + +impl Display for ListContainsOptions { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + if self.sql_null_semantics { + write!(f, "sql_null_semantics")?; + } + Ok(()) + } +} + +impl ListContainsOptions { + /// The declared nullability of a `list_contains` result: a null list or needle can yield + /// null, and under [`sql_null_semantics`](Self::sql_null_semantics) so can a null element. + pub fn result_nullability(&self, list_dtype: &DType, needle_dtype: &DType) -> Nullability { + let mut nullability = list_dtype.nullability().bitor(needle_dtype.nullability()); + if self.sql_null_semantics + && let Some(element_dtype) = list_dtype.as_any_size_list_element_opt() + { + nullability |= element_dtype.nullability(); + } + nullability + } +} + impl ListContains { - /// Creates a lazy list membership check for `needle` in `list`. + /// Creates a lazy list membership check for `needle` in `list`, with a null element that + /// never matches. /// /// # Errors /// - /// Returns an error if the children have different lengths or `list` is not a list array. + /// Returns an error if the children have different lengths, `list` is not a list array, + /// or its element dtype differs from the needle's dtype ignoring nullability. pub fn try_new(list: ArrayRef, needle: ArrayRef) -> VortexResult { - ScalarFnArray::try_new(ListContains.bind(EmptyOptions), vec![list, needle]) + Self::try_new_opts(list, needle, ListContainsOptions::default()) + } + + /// Creates a lazy list membership check for `needle` in `list` with explicit + /// [`ListContainsOptions`]. + /// + /// # Errors + /// + /// Returns an error if the children have different lengths, `list` is not a list array, + /// or its element dtype differs from the needle's dtype ignoring nullability. + pub fn try_new_opts( + list: ArrayRef, + needle: ArrayRef, + options: ListContainsOptions, + ) -> VortexResult { + ScalarFnArray::try_new(ListContains.bind(options), vec![list, needle]) } } impl ScalarFnVTable for ListContains { - type Options = EmptyOptions; + type Options = ListContainsOptions; fn id(&self) -> ScalarFnId { static ID: CachedId = CachedId::new("vortex.list.contains"); *ID } - fn serialize(&self, _instance: &Self::Options) -> VortexResult>> { - Ok(Some(vec![])) + fn serialize(&self, instance: &Self::Options) -> VortexResult>> { + // The default encodes to no bytes, which is also what every file written before the + // option existed carries. + Ok(Some( + pb::ListContainsOpts { + sql_null_semantics: instance.sql_null_semantics, + } + .encode_to_vec(), + )) } fn deserialize( &self, - _metadata: &[u8], + metadata: &[u8], _session: &VortexSession, ) -> VortexResult { - Ok(EmptyOptions) + let opts = pb::ListContainsOpts::decode(metadata)?; + Ok(ListContainsOptions { + sql_null_semantics: opts.sql_null_semantics, + }) } fn arity(&self, _options: &Self::Options) -> Arity { @@ -96,27 +166,32 @@ impl ScalarFnVTable for ListContains { ), } } - fn return_dtype(&self, _options: &Self::Options, arg_dtypes: &[DType]) -> VortexResult { + fn return_dtype(&self, options: &Self::Options, arg_dtypes: &[DType]) -> VortexResult { let list_dtype = &arg_dtypes[0]; let needle_dtype = &arg_dtypes[1]; - let nullability = match list_dtype { - DType::List(_, list_nullability) => list_nullability, - _ => { - vortex_bail!( - "First argument to ListContains must be a List, got {:?}", - list_dtype - ); - } + let DType::List(element_dtype, _) = list_dtype else { + vortex_bail!( + "First argument to ListContains must be a List, got {:?}", + list_dtype + ); + }; + if !element_dtype.eq_ignore_nullability(needle_dtype) { + vortex_bail!( + "Element type {} of list does not match search value {}", + element_dtype, + needle_dtype, + ); } - .bitor(needle_dtype.nullability()); - Ok(DType::Bool(nullability)) + Ok(DType::Bool( + options.result_nullability(list_dtype, needle_dtype), + )) } fn execute( &self, - _options: &Self::Options, + options: &Self::Options, args: &dyn ExecutionArgs, ctx: &mut ExecutionCtx, ) -> VortexResult { @@ -126,16 +201,17 @@ impl ScalarFnVTable for ListContains { if let Some(list_scalar) = list_array.as_constant() && let Some(value_scalar) = value_array.as_constant() { - let result = compute_contains_scalar(&list_scalar, &value_scalar)?; + let result = compute_contains_scalar(&list_scalar, &value_scalar, options)?; return Ok(ConstantArray::new(result, args.row_count()).into_array()); } - compute_list_contains(&list_array, &value_array, ctx) + compute_list_contains(&list_array, &value_array, options, ctx) } - // An empty list can produce false even when the needle is null. - fn is_strict(&self, _options: &Self::Options) -> bool { - false + // Off SQL null semantics an empty list answers `false` even for a null needle; on them a null + // needle is null against any list, as a null list always is. + fn is_strict(&self, options: &Self::Options) -> bool { + options.sql_null_semantics } fn is_infallible(&self, _options: &Self::Options) -> bool { @@ -143,8 +219,18 @@ impl ScalarFnVTable for ListContains { } } -fn compute_contains_scalar(list: &Scalar, needle: &Scalar) -> VortexResult { - let nullability = list.dtype().nullability() | needle.dtype().nullability(); +fn compute_contains_scalar( + list: &Scalar, + needle: &Scalar, + options: &ListContainsOptions, +) -> VortexResult { + if !matches!(list.dtype(), DType::List(..)) { + vortex_bail!( + "First argument to ListContains must be a List, got {}", + list.dtype() + ); + } + let nullability = options.result_nullability(list.dtype(), needle.dtype()); if list.is_null() { return Ok(Scalar::null(DType::Bool(nullability))); @@ -155,21 +241,25 @@ fn compute_contains_scalar(list: &Scalar, needle: &Scalar) -> VortexResult VortexResult { let DType::List(elem_dtype, _) = array.dtype() else { @@ -183,7 +273,9 @@ fn compute_list_contains( ); } - if value.all_invalid(ctx)? || array.all_invalid(ctx)? { + // A null needle is null against every list only when the function is strict: otherwise an + // empty list still answers `false`, which the paths below work out per list. + if array.all_invalid(ctx)? || (options.sql_null_semantics && value.all_invalid(ctx)?) { return Ok(ConstantArray::new( Scalar::null(DType::Bool(Nullability::Nullable)), array.len(), @@ -191,15 +283,23 @@ fn compute_list_contains( .into_array()); } - let nullability = array.dtype().nullability() | value.dtype().nullability(); + let nullability = options.result_nullability(array.dtype(), value.dtype()); if let Some(value_scalar) = value.as_constant() { - list_contains_scalar(array, &value_scalar, nullability, ctx) - } else if let Some(list_scalar) = array.as_constant() { - constant_list_scalar_contains(&list_scalar.as_list(), value, nullability) - } else { - todo!("unsupported list contains with list and element as arrays") + return list_contains_scalar(array, &value_scalar, nullability, options, ctx); } + + if let Some(list_scalar) = array.as_constant() { + return constant_list_scalar_contains( + &list_scalar.as_list(), + value, + nullability, + options, + ctx, + ); + } + + todo!("unsupported list contains with list and element as arrays") } /// There is a constant list scalar (haystack) being compared to an array of needles. @@ -207,45 +307,71 @@ fn constant_list_scalar_contains( list_scalar: &ListScalar, values: &ArrayRef, nullability: Nullability, + options: &ListContainsOptions, + ctx: &mut ExecutionCtx, ) -> VortexResult { let elements = list_scalar.elements().vortex_expect("non null"); - let len = values.len(); let false_scalar = Scalar::bool(false, nullability); let result = elements .iter() .map(|element| { - Binary::try_new( + let comparison = Binary::try_new( ConstantArray::new(element.clone(), len).into_array(), values.clone(), Operator::Eq, )? - .into_array() - .fill_null(false_scalar.clone()) + .into_array(); + if options.sql_null_semantics { + Ok(comparison) + } else { + comparison.fill_null(false_scalar.clone()) + } }) .collect::>>()? .into_iter() .try_reduce_balanced(|acc, res| acc.binary(res, Operator::Or))?; - Ok(result.unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array())) + let matches = result + .unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()) + .execute::(ctx)?; + let validity = if elements.is_empty() && !options.sql_null_semantics { + Validity::NonNullable + } else if options.sql_null_semantics { + matches.validity()?.and(values.validity()?)? + } else { + values.validity()? + }; + Ok(BoolArray::new( + matches.to_bit_buffer(), + validity.union_nullability(nullability), + ) + .into_array()) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. +/// +/// This is the canonical implementation, over an executed [`ListViewArray`]. fn list_contains_scalar( array: &ArrayRef, value: &Scalar, nullability: Nullability, + options: &ListContainsOptions, ctx: &mut ExecutionCtx, ) -> VortexResult { // If the list array is constant, we perform a single comparison. if array.len() > 1 && array.is::() { - let contains = list_contains_scalar(&array.slice(0..1)?, value, nullability, ctx)?; + let contains = list_contains_scalar(&array.slice(0..1)?, value, nullability, options, ctx)?; return Ok(ConstantArray::new(contains.execute_scalar(0, ctx)?, array.len()).into_array()); } let list_array = array.clone().execute::(ctx)?; + if value.is_null() { + return null_needle_in_lists(&list_array, nullability, ctx); + } + let elems = list_array.elements(); if elems.is_empty() { // Must return false when a list is empty (but valid), or null when the list itself is null. @@ -257,7 +383,39 @@ fn list_contains_scalar( Binary::try_new(elems.clone(), rhs.clone().into_array(), Operator::Eq)?.into_array(); // TODO(ngates): we should execute this into a Columnar and check for constant. - let matches = matching_elements.execute::(ctx)?; + let mut matches = matching_elements.execute::(ctx)?; + let valid = matches.validity()?.execute_mask(matches.len(), ctx)?; + + // Under SQL null semantics a list holding a null element answers `null` for a needle that + // matches none of its other elements. The needle is non-null here, so a null comparison is a + // null element; fold "any match" and "any null" per list and keep only the decided rows. + if options.sql_null_semantics && !valid.all_true() { + let valid = valid.to_bit_buffer(); + let any_true = fold_lists( + BoolArray::new(&matches.to_bit_buffer() & &valid, Validity::NonNullable), + &list_array, + ctx, + )?; + let any_null = fold_lists( + BoolArray::new(!valid, Validity::NonNullable), + &list_array, + ctx, + )?; + let decided = &any_true | &!any_null; + let validity = list_array + .validity()? + .and(Validity::from(decided))? + .union_nullability(nullability); + return Ok(BoolArray::new(any_true, validity).into_array()); + } + + // Null comparisons carry unspecified value bits, which must not contribute a match. + if !valid.all_true() { + matches = BoolArray::new( + &matches.to_bit_buffer() & &valid.to_bit_buffer(), + Validity::NonNullable, + ); + } // Fast path: no elements match. if let Some(pred) = matches.as_constant() { @@ -273,13 +431,7 @@ fn list_contains_scalar( list_false_or_null(&list_array, nullability, ctx) } // No elements match, and all comparisons are valid (result in `false`). - Some(false) => { - // False, but match the nullability to the input list array. - Ok( - ConstantArray::new(Scalar::bool(false, nullability), list_array.len()) - .into_array(), - ) - } + Some(false) => list_false_or_null(&list_array, nullability, ctx), // All elements match, and all comparisons are valid (result in `true`). Some(true) => { // True, unless the list itself is empty or NULL. @@ -288,6 +440,21 @@ fn list_contains_scalar( }; } + let list_matches = fold_lists(matches, &list_array, ctx)?; + + Ok(BoolArray::new( + list_matches, + list_array.validity()?.union_nullability(nullability), + ) + .into_array()) +} + +/// For each list, whether any set bit of `matches` falls in the list's element range. +fn fold_lists( + matches: BoolArray, + list_array: &ListViewArray, + ctx: &mut ExecutionCtx, +) -> VortexResult { // Get the offsets and sizes as primitive arrays. They are non-negative, so reinterpret to // unsigned and dispatch over the 4 unsigned widths each (4x4 instead of 8x8). let offsets = list_array @@ -298,18 +465,11 @@ fn list_contains_scalar( let sizes = list_array.sizes().clone().execute::(ctx)?; let sizes = sizes.reinterpret_cast(sizes.ptype().to_unsigned()); - // Process based on the offset and size types. - let list_matches = match_each_unsigned_integer_ptype!(offsets.ptype(), |O| { + Ok(match_each_unsigned_integer_ptype!(offsets.ptype(), |O| { match_each_unsigned_integer_ptype!(sizes.ptype(), |S| { process_matches::(matches, list_array.len(), offsets, sizes, ctx) }) - }); - - Ok(BoolArray::new( - list_matches, - list_array.validity()?.union_nullability(nullability), - ) - .into_array()) + })) } /// Returns a [`BitBuffer`] where each bit represents if a list contains the scalar, derived from a @@ -337,7 +497,7 @@ where // BitIndexIterator yields indices of true bits only. If `.next()` returns // `Some(_)`, at least one element in this list's range matches. - let mut set_bits = BitIndexIterator::new(bits.inner(), offset, size); + let mut set_bits = BitIndexIterator::new(bits.inner(), bits.offset() + offset, size); set_bits.next().is_some() }, ctx.allocator().clone(), @@ -379,6 +539,38 @@ fn list_false_or_null( } } +/// A null needle against each list, off SQL null semantics: `false` for an empty list, which holds +/// nothing to compare it to, and `null` for any other list. +fn null_needle_in_lists( + list_array: &ListViewArray, + nullability: Nullability, + ctx: &mut ExecutionCtx, +) -> VortexResult { + let empty = !&non_empty_lists(list_array, ctx)?; + let validity = list_array + .validity()? + .and(Validity::from(empty))? + .union_nullability(nullability); + Ok(BoolArray::new( + BitBuffer::new_unset_in(list_array.len(), ctx.allocator().clone()), + validity, + ) + .into_array()) +} + +/// One bit per list, set when the list holds at least one element. +fn non_empty_lists(list_array: &ListViewArray, ctx: &mut ExecutionCtx) -> VortexResult { + let sizes = list_array.sizes().clone().execute::(ctx)?; + Ok(match_each_integer_ptype!(sizes.ptype(), |S| { + let sizes = sizes.as_slice::(); + BitBuffer::collect_bool_in( + sizes.len(), + |idx| sizes[idx] != S::zero(), + ctx.allocator().clone(), + ) + })) +} + /// Returns a `Bool` array with `true` for lists which are NOT empty, or `false` if they are empty, /// or `NULL` if the list itself is null. fn list_is_not_empty( @@ -395,19 +587,9 @@ fn list_is_not_empty( .into_array()); } - let sizes = list_array.sizes().clone().execute::(ctx)?; - let buffer = match_each_integer_ptype!(sizes.ptype(), |S| { - let sizes = sizes.as_slice::(); - BitBuffer::collect_bool_in( - sizes.len(), - |idx| sizes[idx] != S::zero(), - ctx.allocator().clone(), - ) - }); - // Copy over the validity mask from the input. Ok(BoolArray::new( - buffer, + non_empty_lists(list_array, ctx)?, list_array.validity()?.union_nullability(nullability), ) .into_array()) @@ -422,15 +604,21 @@ mod tests { use rstest::rstest; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; + use vortex_buffer::buffer; use vortex_error::VortexExpect; use vortex_error::VortexResult; use vortex_session::VortexSession; + use super::process_matches; use crate::ArrayRef; use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; + use crate::arrays::BoolArray; + use crate::arrays::ConstantArray; use crate::arrays::ListArray; + use crate::arrays::ListViewArray; + use crate::arrays::PrimitiveArray; use crate::arrays::VarBinArray; use crate::assert_arrays_eq; use crate::dtype::DType; @@ -442,17 +630,17 @@ mod tests { use crate::expr::col; use crate::expr::get_item; use crate::expr::gt; + use crate::expr::in_list; use crate::expr::list_contains; + use crate::expr::list_contains_opts; use crate::expr::lit; use crate::expr::lt; use crate::expr::or; use crate::expr::root; use crate::expr::stats::Stat; use crate::scalar::Scalar; - use crate::scalar_fn::fns::list_contains::BoolArray; - use crate::scalar_fn::fns::list_contains::ConstantArray; - use crate::scalar_fn::fns::list_contains::ListViewArray; - use crate::scalar_fn::fns::list_contains::PrimitiveArray; + use crate::scalar_fn::fns::list_contains::ListContains; + use crate::scalar_fn::fns::list_contains::ListContainsOptions; use crate::stats::StatsSession; use crate::stats::stat as stat_expr; use crate::validity::Validity; @@ -790,10 +978,12 @@ mod tests { Some("a"), bool_array(vec![false, false, false], Validity::NonNullable) )] + // A null needle is null against a list with elements, but an empty list holds nothing to + // compare it to. #[case( null_strings(vec![vec![], vec![None, None], vec![None, None, None]]), None, - bool_array(vec![false, true, true], Validity::AllInvalid) + bool_array(vec![false, false, false], Validity::from_iter([true, false, false])) )] #[case( null_strings(vec![vec![], vec![None, None], vec![None, None, None]]), @@ -985,4 +1175,317 @@ mod tests { let expected_zero = BoolArray::from_iter([true, false, false, false]); assert_arrays_eq!(result_zero, expected_zero, &mut ctx); } + + const SQL: ListContainsOptions = ListContainsOptions { + sql_null_semantics: true, + }; + + #[rstest] + #[case::default(ListContainsOptions::default(), Some(false))] + #[case::sql(SQL, None)] + fn test_constant_needle_ignores_null_value_bits( + #[case] options: ListContainsOptions, + #[case] null_element_result: Option, + ) -> VortexResult<()> { + // Both the null element and the valid match have the same physical value. + let elements = PrimitiveArray::new( + buffer![7i32, 7, 8], + Validity::from(BitBuffer::from_iter([false, true, true])), + ); + let lists = ListArray::try_new( + elements.into_array(), + buffer![0u32, 1, 2, 3, 3].into_array(), + Validity::AllValid, + )? + .into_array(); + assert_result( + lists.apply(&list_contains_opts(root(), lit(7i32), options)), + [null_element_result, Some(true), Some(false), Some(false)], + ) + } + + #[rstest] + #[case::default(ListContainsOptions::default())] + #[case::sql(SQL)] + fn test_no_matches_preserves_null_lists( + #[case] options: ListContainsOptions, + ) -> VortexResult<()> { + let lists = ListArray::try_new( + buffer![1i32, 2].into_array(), + buffer![0u32, 1, 2].into_array(), + Validity::from(BitBuffer::from_iter([true, false])), + )? + .into_array(); + assert_result( + lists.apply(&list_contains_opts(root(), lit(9i32), options)), + [Some(false), None], + ) + } + + #[test] + fn test_fold_sliced_match_bitmap() { + let mut ctx = array_session().create_execution_ctx(); + let bits = BitBuffer::from_iter([true, false, true, false]).slice(1..); + let result = process_matches::( + BoolArray::new(bits, Validity::NonNullable), + 3, + PrimitiveArray::from_iter([0u32, 1, 2]), + PrimitiveArray::from_iter([1u32, 1, 1]), + &mut ctx, + ); + assert_eq!(result, BitBuffer::from_iter([false, true, false])); + } + + #[rstest] + #[case::default(ListContainsOptions::default())] + #[case::sql(SQL)] + fn test_reject_mismatched_needle_type(#[case] options: ListContainsOptions) { + let list = lit(i32_set(vec![Some(1)])); + let dtype = DType::Utf8(Nullability::NonNullable); + assert!( + list_contains_opts(list.clone(), root(), options) + .bind(&dtype) + .is_err() + ); + assert!( + list_contains_opts(list, lit("1"), options) + .bind(&dtype) + .is_err() + ); + } + + /// A constant `List` set, so that it can hold a null element. + fn i32_set(values: Vec>) -> Scalar { + let element = DType::Primitive(I32, Nullability::Nullable); + Scalar::list( + Arc::new(element.clone()), + values + .into_iter() + .map(|v| match v { + Some(v) => Scalar::primitive(v, Nullability::Nullable), + None => Scalar::null(element.clone()), + }) + .collect(), + Nullability::NonNullable, + ) + } + + fn assert_result( + result: VortexResult, + expected: impl IntoIterator>, + ) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + assert_arrays_eq!(result?, BoolArray::from_iter(expected), &mut ctx); + Ok(()) + } + + #[test] + fn constant_set_null_element_never_matches_by_default() -> VortexResult<()> { + let needles = + PrimitiveArray::from_option_iter([Some(1), Some(2), None, Some(4)]).into_array(); + let set = i32_set(vec![Some(2), None, Some(4)]); + assert_result( + needles.apply(&list_contains(lit(set), root())), + [Some(false), Some(true), None, Some(true)], + ) + } + + #[test] + fn constant_set_null_element_is_unknown_under_sql_semantics() -> VortexResult<()> { + // `x IN (2, NULL, 4)`: a match is still true, but a non-match is null, so `NOT IN` never + // admits it. A null needle stays null either way. + let needles = + PrimitiveArray::from_option_iter([Some(1), Some(2), None, Some(4)]).into_array(); + let set = i32_set(vec![Some(2), None, Some(4)]); + assert_result( + needles.clone().apply(&in_list(root(), lit(set.clone()))), + [None, Some(true), None, Some(true)], + )?; + assert_result( + needles.apply(&crate::expr::not(in_list(root(), lit(set)))), + [None, Some(false), None, Some(false)], + ) + } + + #[test] + fn sql_semantics_without_a_null_element_match_the_default() -> VortexResult<()> { + let needles = PrimitiveArray::from_option_iter([Some(1), Some(2), None]).into_array(); + let set = i32_set(vec![Some(2)]); + assert_result( + needles.apply(&in_list(root(), lit(set))), + [Some(false), Some(true), None], + ) + } + + fn empty_i32_list() -> Scalar { + Scalar::list( + Arc::new(DType::Primitive(I32, Nullability::Nullable)), + vec![], + Nullability::NonNullable, + ) + } + + #[rstest] + #[case::default(ListContainsOptions::default(), Some(false))] + #[case::sql(SQL, None)] + fn null_needle_in_an_empty_list( + #[case] options: ListContainsOptions, + #[case] expected: Option, + ) -> VortexResult<()> { + // Off SQL null semantics an empty list holds nothing to compare a null needle to; on them + // the function is strict. Every evaluation path has to agree, a single row included. + let mut ctx = array_session().create_execution_ctx(); + let null = Scalar::null(DType::Primitive(I32, Nullability::Nullable)); + + let scalar = ConstantArray::new(null, 2) + .into_array() + .apply(&list_contains_opts(lit(empty_i32_list()), root(), options))?; + assert_eq!( + scalar.execute_scalar(0, &mut ctx)?.as_bool().value(), + expected + ); + + let needles = PrimitiveArray::from_option_iter([Some(1), None]).into_array(); + let column = needles.apply(&list_contains_opts(lit(empty_i32_list()), root(), options))?; + assert_eq!( + column.execute_scalar(1, &mut ctx)?.as_bool().value(), + expected + ); + assert_arrays_eq!( + column, + BoolArray::from_iter([Some(false), expected]), + &mut ctx + ); + Ok(()) + } + + #[rstest] + #[case::default(ListContainsOptions::default(), [None, Some(false), None])] + #[case::sql(SQL, [None, None, None])] + fn null_needle_in_a_list_column( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 3], + ) -> VortexResult<()> { + // Lists `[1]`, `[]` and a null list against a null needle. + let lists = ListArray::try_new( + PrimitiveArray::from_option_iter([Some(1i32)]).into_array(), + PrimitiveArray::from_iter(vec![0, 1, 1, 1]).into_array(), + Validity::from_iter([true, true, false]), + )? + .into_array(); + let null = Scalar::null(DType::Primitive(I32, Nullability::Nullable)); + assert_result( + lists.apply(&list_contains_opts(root(), lit(null), options)), + expected, + ) + } + + #[test] + fn strict_only_under_sql_semantics() { + let default = list_contains(lit(empty_i32_list()), root()); + let sql = in_list(root(), lit(empty_i32_list())); + assert!( + !default + .as_scalar() + .is_some_and(|f| f.signature().is_strict()) + ); + assert!(sql.as_scalar().is_some_and(|f| f.signature().is_strict())); + } + + #[test] + fn constant_set_of_only_nulls() -> VortexResult<()> { + // Every element is dropped from the set, so nothing matches; under SQL null semantics no + // answer is known instead. + let needles = PrimitiveArray::from_option_iter([Some(1), None]).into_array(); + let set = i32_set(vec![None, None]); + assert_result( + needles + .clone() + .apply(&list_contains(lit(set.clone()), root())), + [Some(false), None], + )?; + assert_result(needles.apply(&in_list(root(), lit(set))), [None, None]) + } + + #[test] + fn list_array_null_elements_under_sql_semantics() -> VortexResult<()> { + // Lists `[1, null]`, `[2]`, `[3]`, `[]` against a constant needle. + let lists = ListArray::try_new( + PrimitiveArray::from_option_iter([Some(1), None, Some(2), Some(3)]).into_array(), + PrimitiveArray::from_iter(vec![0, 2, 3, 4, 4]).into_array(), + Validity::AllValid, + )? + .into_array(); + assert_result( + lists.clone().apply(&list_contains(root(), lit(5))), + [Some(false), Some(false), Some(false), Some(false)], + )?; + assert_result( + lists + .clone() + .apply(&list_contains_opts(root(), lit(5), SQL)), + [None, Some(false), Some(false), Some(false)], + )?; + assert_result( + lists.apply(&list_contains_opts(root(), lit(1), SQL)), + [Some(true), Some(false), Some(false), Some(false)], + ) + } + + #[test] + fn scalar_needle_in_scalar_set_under_sql_semantics() -> VortexResult<()> { + let set = i32_set(vec![Some(2), None]); + let mut ctx = array_session().create_execution_ctx(); + let hit = ConstantArray::new(Scalar::primitive(2i32, Nullability::Nullable), 1) + .into_array() + .apply(&in_list(root(), lit(set.clone())))?; + assert_eq!( + hit.execute_scalar(0, &mut ctx)?.as_bool().value(), + Some(true) + ); + let miss = ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 1) + .into_array() + .apply(&in_list(root(), lit(set)))?; + assert!(miss.execute_scalar(0, &mut ctx)?.is_null()); + Ok(()) + } + + #[test] + fn sql_semantics_widen_the_return_dtype_to_the_element_nullability() -> VortexResult<()> { + let dtype = DType::Struct( + StructFields::new( + ["a"].into(), + vec![DType::Primitive(I32, Nullability::NonNullable)], + ), + Nullability::NonNullable, + ); + let set = lit(i32_set(vec![Some(1)])); + assert_eq!( + list_contains(set.clone(), col("a")).return_dtype(&dtype)?, + DType::Bool(Nullability::NonNullable) + ); + assert_eq!( + in_list(col("a"), set).return_dtype(&dtype)?, + DType::Bool(Nullability::Nullable) + ); + Ok(()) + } + + #[test] + fn options_round_trip_through_serde() -> VortexResult<()> { + use crate::scalar_fn::ScalarFnVTable; + let session = array_session(); + for options in [ListContainsOptions::default(), SQL] { + let bytes = ListContains + .serialize(&options)? + .vortex_expect("serialized"); + assert_eq!(ListContains.deserialize(&bytes, &session)?, options); + } + // Files written before the option existed carry no bytes at all. + assert_eq!( + ListContains.deserialize(&[], &session)?, + ListContainsOptions::default() + ); + Ok(()) + } } diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 2ab975ecfd7..8e6f29f28a1 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -30,9 +30,9 @@ use vortex::expr::and_collect; use vortex::expr::byte_length; use vortex::expr::cast; use vortex::expr::get_item; +use vortex::expr::in_list; use vortex::expr::is_not_null; use vortex::expr::is_null; -use vortex::expr::list_contains; use vortex::expr::list_length; use vortex::expr::lit; use vortex::expr::nested_case_when; @@ -371,9 +371,9 @@ impl ExpressionConvertor for DefaultExpressionConvertor { return Ok(is_not_null(arg)); } - if let Some(in_list) = df.downcast_ref::() { - let value = self.convert(in_list.expr().as_ref())?; - let list_elements: Vec<_> = in_list + if let Some(in_list_expr) = df.downcast_ref::() { + let value = self.convert(in_list_expr.expr().as_ref())?; + let list_elements: Vec<_> = in_list_expr .list() .iter() .map(|e| { @@ -385,14 +385,26 @@ impl ExpressionConvertor for DefaultExpressionConvertor { }) .try_collect()?; - let list = Scalar::list( - list_elements[0].dtype().clone(), - list_elements, - Nullability::Nullable, - ); - let expr = list_contains(lit(list), value); + let element_dtype = list_elements + .first() + .ok_or_else(|| exec_datafusion_err!("Cannot infer the type of an empty IN list"))? + .dtype() + .with_nullability(list_elements.iter().fold( + Nullability::NonNullable, + |nullability, element| nullability | element.dtype().nullability(), + )); + let list_elements = list_elements + .iter() + .map(|element| { + element + .cast(&element_dtype) + .map_err(|e| exec_datafusion_err!("Invalid IN list element: {e}")) + }) + .collect::>>()?; + let list = Scalar::list(element_dtype, list_elements, Nullability::NonNullable); + let expr = in_list(value, lit(list)); - return Ok(if in_list.negated() { not(expr) } else { expr }); + return Ok(if in_list_expr.negated() { not(expr) } else { expr }); } if let Some(scalar_fn) = df.downcast_ref::() { @@ -727,6 +739,8 @@ mod tests { use arrow_schema::Schema; use arrow_schema::TimeUnit as ArrowTimeUnit; use datafusion::arrow::array::AsArray; + use datafusion::arrow::array::Int32Array; + use datafusion::arrow::array::RecordBatch; use datafusion::arrow::datatypes::Int32Type; use datafusion_common::ScalarValue; use datafusion_common::config::ConfigOptions; @@ -736,6 +750,9 @@ mod tests { use datafusion_physical_plan::expressions as df_expr; use insta::assert_snapshot; use rstest::rstest; + use vortex::array::VortexSessionExecute; + use vortex::array::arrays::BoolArray; + use vortex::array::assert_arrays_eq; use super::*; use crate::common_tests::TestSessionContext; @@ -872,6 +889,65 @@ mod tests { assert_snapshot!(result.display_tree().to_string(), @"vortex.literal(42i32)"); } + #[rstest] + fn test_in_list_null_semantics( + #[values(false, true)] negated: bool, + #[values(false, true)] null_first: bool, + ) { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Int32, true)])); + let batch = RecordBatch::try_new( + schema.clone(), + vec![Arc::new(Int32Array::from(vec![Some(1), Some(2), None]))], + ) + .unwrap(); + let mut list = vec![ + Arc::new(df_expr::Literal::new(ScalarValue::Int32(Some(1)))) as Arc, + Arc::new(df_expr::Literal::new(ScalarValue::Int32(None))) as Arc, + ]; + if null_first { + list.reverse(); + } + let expr = df_expr::InListExpr::try_new( + Arc::new(df_expr::Column::new("value", 0)), + list, + negated, + &schema, + ) + .unwrap(); + let expected = expr + .evaluate(&batch) + .unwrap() + .into_array(batch.num_rows()) + .unwrap(); + let session = VortexSession::default(); + let converted = DefaultExpressionConvertor::new(session.clone()) + .convert(&expr) + .unwrap(); + let input = session + .arrow() + .from_arrow_record_batch(batch, &schema) + .unwrap(); + let actual = input.apply(&converted).unwrap(); + assert_arrays_eq!( + actual, + BoolArray::from_iter(expected.as_boolean().iter()), + &mut session.create_execution_ctx(), + ); + } + + #[test] + fn test_empty_in_list_declines_conversion() { + let schema = Schema::new(vec![Field::new("value", DataType::Int32, true)]); + let expr = df_expr::InListExpr::try_new_from_array( + Arc::new(df_expr::Column::new("value", 0)), + Arc::new(Int32Array::from(Vec::::new())), + false, + &schema, + ) + .unwrap(); + assert!(DefaultExpressionConvertor::default().convert(&expr).is_err()); + } + #[test] fn test_expr_from_df_binary() { let left = Arc::new(df_expr::Column::new("left", 0)) as Arc; From 146cd74e78e17fef5dfa6ffd622feccc2fd3fe1b Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 10:59:15 -0400 Subject: [PATCH 2/6] Fix DataFusion membership test compilation Remove the unsupported trailing macro comma and clone the schema Arc explicitly to satisfy the repository lint configuration. Signed-off-by: Robert Kruszewski --- vortex-datafusion/src/convert/exprs.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 8e6f29f28a1..a0c7a85dbbc 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -896,7 +896,7 @@ mod tests { ) { let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Int32, true)])); let batch = RecordBatch::try_new( - schema.clone(), + Arc::clone(&schema), vec![Arc::new(Int32Array::from(vec![Some(1), Some(2), None]))], ) .unwrap(); @@ -931,7 +931,7 @@ mod tests { assert_arrays_eq!( actual, BoolArray::from_iter(expected.as_boolean().iter()), - &mut session.create_execution_ctx(), + &mut session.create_execution_ctx() ); } From 5c10f3bc4a227074a94bf0f64ffbb32eb4c06761 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 12:27:36 -0400 Subject: [PATCH 3/6] Format SQL membership conversion Signed-off-by: Robert Kruszewski --- vortex-datafusion/src/convert/exprs.rs | 29 +++++++++++++++++++------- 1 file changed, 22 insertions(+), 7 deletions(-) diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index a0c7a85dbbc..1581f9f83ce 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -389,10 +389,13 @@ impl ExpressionConvertor for DefaultExpressionConvertor { .first() .ok_or_else(|| exec_datafusion_err!("Cannot infer the type of an empty IN list"))? .dtype() - .with_nullability(list_elements.iter().fold( - Nullability::NonNullable, - |nullability, element| nullability | element.dtype().nullability(), - )); + .with_nullability( + list_elements + .iter() + .fold(Nullability::NonNullable, |nullability, element| { + nullability | element.dtype().nullability() + }), + ); let list_elements = list_elements .iter() .map(|element| { @@ -404,7 +407,11 @@ impl ExpressionConvertor for DefaultExpressionConvertor { let list = Scalar::list(element_dtype, list_elements, Nullability::NonNullable); let expr = in_list(value, lit(list)); - return Ok(if in_list_expr.negated() { not(expr) } else { expr }); + return Ok(if in_list_expr.negated() { + not(expr) + } else { + expr + }); } if let Some(scalar_fn) = df.downcast_ref::() { @@ -894,7 +901,11 @@ mod tests { #[values(false, true)] negated: bool, #[values(false, true)] null_first: bool, ) { - let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Int32, true)])); + let schema = Arc::new(Schema::new(vec![Field::new( + "value", + DataType::Int32, + true, + )])); let batch = RecordBatch::try_new( Arc::clone(&schema), vec![Arc::new(Int32Array::from(vec![Some(1), Some(2), None]))], @@ -945,7 +956,11 @@ mod tests { &schema, ) .unwrap(); - assert!(DefaultExpressionConvertor::default().convert(&expr).is_err()); + assert!( + DefaultExpressionConvertor::default() + .convert(&expr) + .is_err() + ); } #[test] From 4d0e5ed0eeeab1f946f68e8400cafa5dd729dac4 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 13:10:44 -0400 Subject: [PATCH 4/6] Declare DataFusion array assertion test dependency Enable the vortex-array test harness explicitly for DataFusion tests so standalone test compilation does not depend on features enabled by other workspace packages. Signed-off-by: Robert Kruszewski --- Cargo.lock | 1 + vortex-datafusion/Cargo.toml | 1 + 2 files changed, 2 insertions(+) diff --git a/Cargo.lock b/Cargo.lock index 37de6cd4aa1..abf7b0b83b7 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11066,6 +11066,7 @@ dependencies = [ "tracing", "url", "vortex", + "vortex-array", "vortex-arrow", "vortex-utils", ] diff --git a/vortex-datafusion/Cargo.toml b/vortex-datafusion/Cargo.toml index 49fb22d4f59..cb60afa7c14 100644 --- a/vortex-datafusion/Cargo.toml +++ b/vortex-datafusion/Cargo.toml @@ -48,6 +48,7 @@ rstest = { workspace = true } tempfile = { workspace = true } tokio = { workspace = true, features = ["test-util", "rt-multi-thread", "fs"] } url = { workspace = true } +vortex-array = { workspace = true, features = ["_test-harness"] } [lints] workspace = true From 904d83b03f4b00b38e32df08045da0b9ad67f122 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 14:52:00 -0400 Subject: [PATCH 5/6] simplify Signed-off-by: Robert Kruszewski --- .../sequence/src/compute/list_contains.rs | 40 ++++---- .../src/scalar_fn/fns/list_contains/mod.rs | 92 +++++++++---------- vortex-datafusion/src/convert/exprs.rs | 26 +++--- 3 files changed, 73 insertions(+), 85 deletions(-) diff --git a/encodings/sequence/src/compute/list_contains.rs b/encodings/sequence/src/compute/list_contains.rs index 5b3e620a71c..69aac6c763b 100644 --- a/encodings/sequence/src/compute/list_contains.rs +++ b/encodings/sequence/src/compute/list_contains.rs @@ -9,6 +9,8 @@ use vortex_array::arrays::ConstantArray; use vortex_array::scalar::Scalar; use vortex_array::scalar_fn::fns::list_contains::ListContainsElementReduce; use vortex_array::scalar_fn::fns::list_contains::ListContainsOptions; +use vortex_array::validity::Validity; +use vortex_buffer::BitBufferMut; use vortex_error::VortexResult; use crate::array::Sequence; @@ -25,25 +27,18 @@ impl ListContainsElementReduce for Sequence { return Ok(None); }; - // A null list scalar has no elements to intersect with. Nothing checks this before the - // reduce rule runs, so fall back to the generic implementation, which resolves a null - // haystack to all-null rather than panicking here. + // A null list falls back to the generic path, which yields all-null. let Some(list_elements) = list_scalar.as_list().elements() else { return Ok(None); }; - // The intersection search treats a null element as matching nothing, which under SQL null - // semantics is only half the answer: there a non-match is unknown, which the search cannot - // express. - if options.sql_null_semantics && list_elements.iter().any(Scalar::is_null) { - return Ok(None); - } - let nullability = options.result_nullability(list.dtype(), element.dtype()); - let mut set_indices: Vec = Vec::new(); + let mut matches = BitBufferMut::new_unset(element.len()); + let mut has_null = false; for intercept in list_elements.iter() { let Some(intercept) = intercept.as_primitive().pvalue() else { + has_null = true; continue; }; match find_intersection( @@ -54,7 +49,7 @@ impl ListContainsElementReduce for Sequence { ) { // Non-integer elements do not match the sequence. None | Some(Intersection::None) => {} - Some(Intersection::At(idx)) => set_indices.push(idx), + Some(Intersection::At(idx)) => matches.set(idx), Some(Intersection::All) => { return Ok(Some( ConstantArray::new(Scalar::bool(true, nullability), element.len()) @@ -64,9 +59,16 @@ impl ListContainsElementReduce for Sequence { } } - Ok(Some( - BoolArray::from_indices(element.len(), set_indices, nullability.into()).into_array(), - )) + let matches = matches.freeze(); + + // Under SQL null semantics a null element makes every non-match unknown. + let validity = if options.sql_null_semantics && has_null { + Validity::from_bit_buffer(matches.clone(), nullability) + } else { + nullability.into() + }; + + Ok(Some(BoolArray::new(matches, validity).into_array())) } } @@ -137,8 +139,7 @@ mod tests { #[test] fn test_list_contains_null_list() { - // A null haystack resolves to null for every row. The reduce rule used to assume the list - // scalar was non-null and panicked instead of declining the reduction. + // A null list yields null for every row. The reduce rule used to panic on it. let array = Sequence::try_new_typed(1, 1, Nullability::NonNullable, 3) .unwrap() .into_array(); @@ -152,9 +153,8 @@ mod tests { #[test] fn test_list_contains_null_element_semantics() { - // The sequence kernel skips a null element, which is right by default. Under SQL null - // semantics a non-match must be null instead, so a constant list holding a null goes to - // the generic path instead. + // By default a null element matches nothing. Under SQL null semantics it makes every + // non-match null. let element = DType::Primitive(I32, Nullability::Nullable); let set = Scalar::list( Arc::new(element.clone()), diff --git a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index d3182d2fdb3..f67536a960b 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -21,6 +21,7 @@ use vortex_session::registry::CachedId; use vortex_utils::iter::ReduceBalancedIterExt; use crate::ArrayRef; +use crate::Columnar; use crate::ExecutionCtx; use crate::IntoArray; use crate::arrays::BoolArray; @@ -290,25 +291,21 @@ fn compute_list_contains( } if let Some(list_scalar) = array.as_constant() { - return constant_list_scalar_contains( - &list_scalar.as_list(), - value, - nullability, - options, - ctx, - ); + return constant_list_scalar_contains(&list_scalar.as_list(), value, nullability, options); } todo!("unsupported list contains with list and element as arrays") } /// There is a constant list scalar (haystack) being compared to an array of needles. +/// +/// The result stays lazy. `Or` is Kleene, so under SQL null semantics the disjunction of the raw +/// comparisons is already the `IN` answer. fn constant_list_scalar_contains( list_scalar: &ListScalar, values: &ArrayRef, nullability: Nullability, options: &ListContainsOptions, - ctx: &mut ExecutionCtx, ) -> VortexResult { let elements = list_scalar.elements().vortex_expect("non null"); let len = values.len(); @@ -333,21 +330,24 @@ fn constant_list_scalar_contains( .into_iter() .try_reduce_balanced(|acc, res| acc.binary(res, Operator::Or))?; - let matches = result - .unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()) - .execute::(ctx)?; - let validity = if elements.is_empty() && !options.sql_null_semantics { - Validity::NonNullable - } else if options.sql_null_semantics { - matches.validity()?.and(values.validity()?)? + let mut result = result.unwrap_or_else(|| ConstantArray::new(false_scalar, len).into_array()); + + // A null needle must still yield null where nothing above keeps it: off SQL null semantics + // `fill_null` erases it, and under them an empty list has no comparison to carry it. + let erases_null_needle = if options.sql_null_semantics { + elements.is_empty() } else { - values.validity()? + !elements.is_empty() }; - Ok(BoolArray::new( - matches.to_bit_buffer(), - validity.union_nullability(nullability), - ) - .into_array()) + if erases_null_needle && values.dtype().is_nullable() { + result = result.mask(values.is_not_null()?)?; + } + + if result.dtype().nullability() != nullability { + result = result.cast(DType::Bool(nullability))?; + } + + Ok(result) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -380,10 +380,23 @@ fn list_contains_scalar( let rhs = ConstantArray::new(value.clone(), elems.len()); let matching_elements = - Binary::try_new(elems.clone(), rhs.clone().into_array(), Operator::Eq)?.into_array(); - - // TODO(ngates): we should execute this into a Columnar and check for constant. - let mut matches = matching_elements.execute::(ctx)?; + Binary::try_new(elems.clone(), rhs.into_array(), Operator::Eq)?.into_array(); + + let mut matches = match matching_elements.execute::(ctx)? { + Columnar::Constant(constant) => { + return match constant.scalar().as_bool().value() { + // The needle is not null, so every comparison is null only if every element is. + None if options.sql_null_semantics => { + null_needle_in_lists(&list_array, nullability, ctx) + } + // No element matches: false, unless the list itself is null. + None | Some(false) => list_false_or_null(&list_array, nullability, ctx), + // Every element matches: true, unless the list itself is empty or null. + Some(true) => list_is_not_empty(&list_array, nullability, ctx), + }; + } + Columnar::Canonical(canonical) => canonical.into_bool(), + }; let valid = matches.validity()?.execute_mask(matches.len(), ctx)?; // Under SQL null semantics a list holding a null element answers `null` for a needle that @@ -417,29 +430,6 @@ fn list_contains_scalar( ); } - // Fast path: no elements match. - if let Some(pred) = matches.as_constant() { - return match pred.as_bool().value() { - // All comparisons are invalid (result in `null`), and search is not null because - // we already checked for null above. - None => { - assert!( - !rhs.scalar().is_null(), - "Search value must not be null here" - ); - // False, unless the list itself is null in which case we return null. - list_false_or_null(&list_array, nullability, ctx) - } - // No elements match, and all comparisons are valid (result in `false`). - Some(false) => list_false_or_null(&list_array, nullability, ctx), - // All elements match, and all comparisons are valid (result in `true`). - Some(true) => { - // True, unless the list itself is empty or NULL. - list_is_not_empty(&list_array, nullability, ctx) - } - }; - } - let list_matches = fold_lists(matches, &list_array, ctx)?; Ok(BoolArray::new( @@ -539,8 +529,10 @@ fn list_false_or_null( } } -/// A null needle against each list, off SQL null semantics: `false` for an empty list, which holds -/// nothing to compare it to, and `null` for any other list. +/// `false` for an empty list, which holds nothing to compare against, and `null` for any other list. +/// +/// This is a null needle off SQL null semantics, or under them a needle against lists that hold +/// only nulls. fn null_needle_in_lists( list_array: &ListViewArray, nullability: Nullability, diff --git a/vortex-datafusion/src/convert/exprs.rs b/vortex-datafusion/src/convert/exprs.rs index 1581f9f83ce..fa2006a515a 100644 --- a/vortex-datafusion/src/convert/exprs.rs +++ b/vortex-datafusion/src/convert/exprs.rs @@ -385,25 +385,21 @@ impl ExpressionConvertor for DefaultExpressionConvertor { }) .try_collect()?; + // A null literal is nullable and a non-null one is not. `Scalar::list` needs one + // element dtype, so make all elements nullable when any is. + let list_elements = if list_elements.iter().any(|e| e.dtype().is_nullable()) { + list_elements + .into_iter() + .map(Scalar::into_nullable) + .collect() + } else { + list_elements + }; let element_dtype = list_elements .first() .ok_or_else(|| exec_datafusion_err!("Cannot infer the type of an empty IN list"))? .dtype() - .with_nullability( - list_elements - .iter() - .fold(Nullability::NonNullable, |nullability, element| { - nullability | element.dtype().nullability() - }), - ); - let list_elements = list_elements - .iter() - .map(|element| { - element - .cast(&element_dtype) - .map_err(|e| exec_datafusion_err!("Invalid IN list element: {e}")) - }) - .collect::>>()?; + .clone(); let list = Scalar::list(element_dtype, list_elements, Nullability::NonNullable); let expr = in_list(value, lit(list)); From de591fc7cbf07b90018c7d21ed4acbd7568a01ea Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 16:42:28 -0400 Subject: [PATCH 6/6] Regenerate protobuf bindings for ListContainsOpts Co-Authored-By: Claude Opus 5.5 Signed-off-by: Robert Kruszewski --- vortex-array/src/proto/generated/vortex.expr.rs | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/vortex-array/src/proto/generated/vortex.expr.rs b/vortex-array/src/proto/generated/vortex.expr.rs index 2aff5a796e7..b08822055a7 100644 --- a/vortex-array/src/proto/generated/vortex.expr.rs +++ b/vortex-array/src/proto/generated/vortex.expr.rs @@ -177,6 +177,11 @@ pub struct LikeOpts { #[prost(bool, tag = "2")] pub case_insensitive: bool, } +#[derive(Clone, Copy, PartialEq, Eq, Hash, ::prost::Message)] +pub struct ListContainsOpts { + #[prost(bool, tag = "1")] + pub sql_null_semantics: bool, +} #[derive(Clone, PartialEq, ::prost::Message)] pub struct CastOpts { #[prost(message, optional, tag = "1")]