From 2d3cfb0f0a905b9e6931a7a68a7a8e12d1374726 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 10:09:14 -0400 Subject: [PATCH 1/9] Optimize list_contains with prepared constant sets Replace comparison/OR chains with prepared membership probes, normalize constant sets, and evaluate list and needle columns together. Include canonicalization support and probe regressions. Build on the separate SQL null-semantics and benchmark changes. Signed-off-by: Robert Kruszewski --- vortex-array/src/arrays/constant/mod.rs | 1 + .../src/arrays/constant/vtable/canonical.rs | 87 ++- vortex-array/src/scalar/typed_view/list.rs | 8 + .../src/scalar_fn/fns/binary/compare/mod.rs | 3 +- .../scalar_fn/fns/binary/compare/nested.rs | 4 +- .../src/scalar_fn/fns/list_contains/kernel.rs | 115 ++- .../src/scalar_fn/fns/list_contains/mod.rs | 692 ++++++++++++++++-- .../fns/list_contains/prepared/mod.rs | 487 ++++++++++++ .../fns/list_contains/prepared/tests.rs | 226 ++++++ 9 files changed, 1534 insertions(+), 89 deletions(-) create mode 100644 vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs create mode 100644 vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs diff --git a/vortex-array/src/arrays/constant/mod.rs b/vortex-array/src/arrays/constant/mod.rs index 36df554d4c3..9bbfb2cc08a 100644 --- a/vortex-array/src/arrays/constant/mod.rs +++ b/vortex-array/src/arrays/constant/mod.rs @@ -15,3 +15,4 @@ pub(crate) mod compute; mod vtable; pub use vtable::Constant; +pub(crate) use vtable::canonical::list_scalar_elements; diff --git a/vortex-array/src/arrays/constant/vtable/canonical.rs b/vortex-array/src/arrays/constant/vtable/canonical.rs index e1f9ad02f6c..c7241a9c753 100644 --- a/vortex-array/src/arrays/constant/vtable/canonical.rs +++ b/vortex-array/src/arrays/constant/vtable/canonical.rs @@ -7,6 +7,7 @@ use itertools::Itertools; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; use vortex_buffer::BufferAllocatorRef; +use vortex_buffer::BufferMut; use vortex_buffer::BufferString; use vortex_buffer::ByteBuffer; use vortex_buffer::buffer; @@ -34,6 +35,8 @@ use crate::arrays::UnionArray; use crate::arrays::VarBinViewArray; use crate::arrays::VariantArray; use crate::arrays::varbinview::BinaryView; +use crate::builders::ArrayBuilder; +use crate::builders::VarBinViewBuilder; use crate::builders::builder_with_capacity_in; use crate::dtype::DType; use crate::dtype::DecimalType; @@ -43,7 +46,9 @@ use crate::match_each_decimal_value_type; use crate::match_each_native_ptype; use crate::match_smallest_list_offset_type; use crate::scalar::DecimalValue; +use crate::scalar::ListScalar; use crate::scalar::Scalar; +use crate::scalar::ScalarValue; use crate::validity::Validity; /// Shared implementation for both `canonicalize` and `execute` methods. @@ -273,25 +278,7 @@ fn constant_canonical_list_array( // Since "canonicalize" only applies to the top level array, we can simply have 1 scalar in our // child `elements` and have all list views point to that scalar. - let elements = if let Some(elements) = list.elements() { - // Extract the list elements out of the scalar into a new array. - let mut builder = builder_with_capacity_in( - list.dtype() - .as_list_element_opt() - .vortex_expect("list scalar somehow did not have a list DType"), - list.len(), - allocator, - ); - for scalar in &elements { - builder - .append_scalar(scalar) - .vortex_expect("list element scalar was invalid"); - } - builder.finish() - } else { - // Otherwise all values are null, and we don't need to store anything in our `elements`. - Canonical::empty(list.element_dtype()).into_array() - }; + let elements = list_scalar_elements(&list, allocator); let validity = if scalar.dtype().is_nullable() { if list.is_null() { @@ -322,6 +309,68 @@ fn constant_canonical_list_array( unsafe { ListViewArray::new_unchecked(elements, offsets, sizes, validity) } } +/// The elements of a list scalar as an array, one row per element; empty for a null list. +/// +/// Primitive and byte elements are written straight from their values rather than through a scalar +/// each. +pub(crate) fn list_scalar_elements(list: &ListScalar, allocator: &BufferAllocatorRef) -> ArrayRef { + let element_dtype = list.element_dtype(); + let Some(values) = list.element_values() else { + return Canonical::empty(element_dtype).into_array(); + }; + + match element_dtype { + DType::Primitive(ptype, nullability) => match_each_native_ptype!(ptype, |T| { + let mut buffer = BufferMut::::with_capacity_in(values.len(), allocator.clone()); + buffer.extend(values.iter().map(|value| { + value.as_ref().map_or_else(T::default, |value| { + value + .as_primitive() + .cast::() + .vortex_expect("list element of the list's element ptype") + }) + })); + PrimitiveArray::new(buffer.freeze(), element_validity(values, *nullability)) + .into_array() + }), + DType::Utf8(_) | DType::Binary(_) => { + let mut builder = VarBinViewBuilder::with_capacity_in( + element_dtype.clone(), + values.len(), + allocator.clone(), + ); + for value in values { + match value { + None => builder.append_null(), + Some(ScalarValue::Utf8(value)) => builder.append_value(value.as_bytes()), + Some(value) => builder.append_value(value.as_binary().as_slice()), + } + } + builder.finish() + } + _ => { + let mut builder = builder_with_capacity_in(element_dtype, values.len(), allocator); + for idx in 0..values.len() { + builder + .append_scalar(&list.element(idx).vortex_expect("index within the list")) + .vortex_expect("list element scalar was invalid"); + } + builder.finish() + } + } +} + +/// The validity of a list's elements, a null element being a `None` value. +fn element_validity(values: &[Option], nullability: Nullability) -> Validity { + match nullability { + Nullability::NonNullable => Validity::NonNullable, + Nullability::Nullable if values.iter().all(Option::is_some) => Validity::AllValid, + Nullability::Nullable => { + Validity::from(BitBuffer::from_iter(values.iter().map(Option::is_some))) + } + } +} + /// Creates a [`FixedSizeListArray`] whose every row holds the same list. fn constant_canonical_fixed_size_list_array( values: Option>, diff --git a/vortex-array/src/scalar/typed_view/list.rs b/vortex-array/src/scalar/typed_view/list.rs index f97857c92ed..359a1b32fcb 100644 --- a/vortex-array/src/scalar/typed_view/list.rs +++ b/vortex-array/src/scalar/typed_view/list.rs @@ -160,6 +160,14 @@ impl<'a> ListScalar<'a> { }) } + /// Returns the values of the list's elements, `None` for a null element, without building a + /// scalar for each. + /// + /// Returns None if the list is null. + pub(crate) fn element_values(&self) -> Option<&'a [Option]> { + self.elements + } + /// Returns all elements in the list as a vector of scalars. /// /// Returns None if the list is null. diff --git a/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs b/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs index a1635c718ff..8a27d18092e 100644 --- a/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs +++ b/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs @@ -48,6 +48,7 @@ mod boolean; mod bytes; mod decimal; mod nested; +pub(crate) use nested::build_comparator as build_row_comparator; mod primitive; #[cfg(test)] mod tests; @@ -304,7 +305,7 @@ pub(super) fn collect_zip_bits( } /// Bit-pack the predicate `f(values[i])` over a slice into a [`BitBuffer`]. -pub(super) fn collect_bits( +pub(crate) fn collect_bits( values: &[T], f: impl Fn(T) -> bool, allocator: &BufferAllocatorRef, diff --git a/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs b/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs index 0ce889c46f1..ccccc3ca955 100644 --- a/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs +++ b/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs @@ -49,7 +49,7 @@ use crate::scalar_fn::fns::binary::compare::compare_validity; use crate::scalar_fn::fns::operators::CompareOperator; /// A row comparator: compares row `i` of the left operand against row `j` of the right operand. -type RowComparator = Box Ordering>; +pub(crate) type RowComparator = Box Ordering>; /// Compare two nested arrays row by row. pub(super) fn compare_nested( @@ -116,7 +116,7 @@ fn validity_mask(array: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult /// Build a row comparator over two recursively canonical arrays of the same logical dtype /// (ignoring nullability). Null values order before all non-null values at every level. -fn build_comparator( +pub(crate) fn build_comparator( lhs: &ArrayRef, rhs: &ArrayRef, ctx: &mut ExecutionCtx, 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 7cee5f7aac1..cab134e6f5f 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -1,30 +1,35 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors +use vortex_buffer::BitBuffer; use vortex_error::VortexExpect; use vortex_error::VortexResult; use crate::ArrayRef; use crate::ExecutionCtx; +use crate::IntoArray; use crate::array::ArrayView; use crate::array::VTable; +use crate::arrays::BoolArray; +use crate::arrays::Constant; use crate::arrays::ScalarFn; +use crate::arrays::constant::list_scalar_elements; use crate::arrays::scalar_fn::ExactScalarFn; use crate::arrays::scalar_fn::ScalarFnArrayExt; use crate::arrays::scalar_fn::ScalarFnArrayView; +use crate::dtype::DType; +use crate::dtype::Nullability; 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; +use crate::validity::Validity; /// Check list-contains without reading buffers (metadata-only). /// /// This trait dispatches on the **element** (needle) child at index 1 of the `ListContains` -/// expression. `Self::Array` is the concrete element encoding, while the list (haystack) is -/// passed as an opaque `&ArrayRef`. -/// -/// A future `ListContainsListReduce` could dispatch on the list side (child 0) for encodings -/// with specialized list representations. +/// expression. `Self` is the concrete element encoding, while the list (haystack) is passed as an +/// opaque `&ArrayRef`. /// /// Return `None` if the operation cannot be resolved from metadata alone. pub trait ListContainsElementReduce: VTable { @@ -49,6 +54,106 @@ pub trait ListContainsElementKernel: VTable { ) -> VortexResult>; } +/// The constant haystack of a `list_contains`, prepared for one pass over a column of needles. +/// +/// This is the `IN` shape, `list_contains(lit([...]), column)`: the list's elements are written +/// once into the array [`elements`](Self::elements), a kernel builds a probe structure from them and +/// tests every needle against it, then hands the membership bits back to [`finish`](Self::finish). +/// Shared by the canonical implementations and the element kernels, so that the result nullability +/// and the SQL null semantics have one definition. +/// +/// A null element is dropped from the elements — it never equals a needle — and only decides +/// whether a non-match is [unknown](Self::non_match_is_unknown). An empty list is an empty set, +/// whose answer for a null needle [`finish`](Self::finish) settles by the options. +/// +/// Preparing the set allocates and executes, so this serves [`ListContainsElementKernel`]. A +/// [`ListContainsElementReduce`] rule, which may not read buffers, works from the list's scalar. +pub struct ListContainsSet { + elements: ArrayRef, + pub(super) nullability: Nullability, + non_match_is_unknown: bool, + /// Whether every needle, a null one included, is absent: an empty list off SQL null semantics. + matches_nothing: bool, +} + +impl ListContainsSet { + /// Prepares the constant list `list` to be probed by needles of dtype `needle_dtype`. + /// + /// Returns `None` when there is no set to probe: a list that is not a constant or is null, or a + /// needle whose dtype does not match the list's elements. + pub fn try_new( + list: &ArrayRef, + needle_dtype: &DType, + options: &ListContainsOptions, + ctx: &mut ExecutionCtx, + ) -> VortexResult> { + let DType::List(element_dtype, _) = list.dtype() else { + return Ok(None); + }; + // A needle of a different dtype is an invalid expression, which the scalar function + // reports as an error. + if !element_dtype.eq_ignore_nullability(needle_dtype) { + return Ok(None); + } + // Every needle is probed against the same set, so the haystack has to be one constant + // list, and a null one holds no elements at all. + let Some(constant) = list.as_opt::() else { + return Ok(None); + }; + let list_scalar = constant.scalar(); + if list_scalar.is_null() { + return Ok(None); + } + + let elements = list_scalar_elements(&list_scalar.as_list(), ctx.allocator()); + let size = elements.len(); + + let valid = elements.validity()?.execute_mask(size, ctx)?; + let non_match_is_unknown = options.sql_null_semantics && !valid.all_true(); + // A null element never equals a needle, so a probe is better off without it. + let elements = if valid.all_true() { + elements + } else { + elements.filter(valid)? + }; + + Ok(Some(Self { + elements, + nullability: options.result_nullability(list.dtype(), needle_dtype), + non_match_is_unknown, + matches_nothing: size == 0 && !options.sql_null_semantics, + })) + } + + /// The list's non-null elements, one row each, in the dtype of the needles. + pub fn elements(&self) -> &ArrayRef { + &self.elements + } + + /// Whether a needle matching no element answers `null` rather than `false`. + /// + /// A null element under [`ListContainsOptions::sql_null_semantics`] makes the comparison + /// unknown, the three-valued semantics of SQL `IN`. + pub fn non_match_is_unknown(&self) -> bool { + self.non_match_is_unknown + } + + /// Assembles the result from one membership bit per needle and the needles' validity. + /// + /// A null needle stays null, unless the list is empty off SQL null semantics, which answers + /// `false` for every needle. When a non-match is unknown, only the rows that matched stay valid. + pub fn finish(&self, bits: BitBuffer, needle_validity: Validity) -> VortexResult { + let validity = if self.matches_nothing { + Validity::NonNullable + } else if self.non_match_is_unknown { + needle_validity.and(Validity::from(bits.clone()))? + } else { + needle_validity + }; + Ok(BoolArray::new(bits, validity.union_nullability(self.nullability)).into_array()) + } +} + /// Adaptor that wraps a [`ListContainsElementReduce`] impl as an [`ArrayParentReduceRule`]. #[derive(Default, Debug)] pub struct ListContainsElementReduceAdaptor(pub V); 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 f67536a960b..e1712d835e9 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -2,23 +2,29 @@ // SPDX-FileCopyrightText: Copyright the Vortex contributors mod kernel; +mod prepared; +use std::cmp::Ordering; use std::fmt::Display; use std::fmt::Formatter; +use std::hash::Hash; +use std::iter; use std::ops::BitOr; +use std::sync::Arc; use arrow_buffer::bit_iterator::BitIndexIterator; pub use kernel::*; use num_traits::Zero; +pub use prepared::PreparedSet; use prost::Message; use vortex_buffer::BitBuffer; -use vortex_error::VortexExpect; +use vortex_buffer::Buffer; +use vortex_buffer::BufferMut; use vortex_error::VortexResult; use vortex_error::vortex_bail; use vortex_error::vortex_err; use vortex_session::VortexSession; use vortex_session::registry::CachedId; -use vortex_utils::iter::ReduceBalancedIterExt; use crate::ArrayRef; use crate::Columnar; @@ -33,14 +39,14 @@ use crate::arrays::ScalarFnArray; use crate::arrays::bool::BoolArrayExt; use crate::arrays::listview::ListViewArraySlotsExt; use crate::arrays::primitive::PrimitiveArrayExt; -use crate::builtins::ArrayBuiltins; use crate::dtype::DType; use crate::dtype::IntegerPType; use crate::dtype::Nullability; +use crate::expr::Expression; +use crate::expr::lit; 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; @@ -49,6 +55,7 @@ use crate::scalar_fn::ScalarFnId; use crate::scalar_fn::ScalarFnVTable; use crate::scalar_fn::ScalarFnVTableExt; use crate::scalar_fn::fns::binary::Binary; +use crate::scalar_fn::fns::literal::Literal; use crate::scalar_fn::fns::operators::Operator; use crate::validity::Validity; @@ -199,16 +206,29 @@ impl ScalarFnVTable for ListContains { let list_array = args.get(0)?; let value_array = args.get(1)?; - if let Some(list_scalar) = list_array.as_constant() + // Borrow the list: a constant list scalar owns every element, so cloning it is not free. + if let Some(list) = list_array.as_opt::() && let Some(value_scalar) = value_array.as_constant() { - let result = compute_contains_scalar(&list_scalar, &value_scalar, options)?; + 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, options, ctx) } + fn simplify_untyped( + &self, + options: &Self::Options, + expr: &Expression, + ) -> VortexResult> { + let Some(list) = expr.child(0).as_opt::() else { + return Ok(None); + }; + Ok(normalized_set(list, options) + .map(|list| ListContains.new_expr(*options, [lit(list), expr.child(1).clone()]))) + } + // 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 { @@ -220,6 +240,56 @@ impl ScalarFnVTable for ListContains { } } +/// A constant list rewritten into the form a set probe wants — its elements sorted, without +/// duplicates, and with at most one null — or `None` when it is in that form already or its +/// elements have no total order. +/// +/// Neither the order of the elements nor their repetition changes a membership test, and one null +/// element decides as much as many. Whether a null survives does matter: under SQL null semantics +/// it makes a non-match unknown, and off them a list of nothing but nulls is still not empty, which +/// a null needle tells apart. +/// +/// Normalizing once, while the expression is optimized, spares every batch the sort. +fn normalized_set(list: &Scalar, options: &ListContainsOptions) -> Option { + let DType::List(element_dtype, nullability) = list.dtype() else { + return None; + }; + if !matches!( + element_dtype.as_ref(), + DType::Bool(_) + | DType::Primitive(..) + | DType::Decimal(..) + | DType::Utf8(_) + | DType::Binary(_) + ) { + return None; + } + let elements = list.as_list().elements()?; + + let had_null = elements.iter().any(Scalar::is_null); + let mut set: Vec = elements + .iter() + .filter(|element| !element.is_null()) + .cloned() + .collect(); + let mut incomparable = false; + set.sort_by(|a, b| { + a.partial_cmp(b).unwrap_or_else(|| { + incomparable = true; + Ordering::Equal + }) + }); + if incomparable { + return None; + } + set.dedup(); + if had_null && (options.sql_null_semantics || set.is_empty()) { + set.push(Scalar::null(element_dtype.as_ref().clone())); + } + + (set != elements).then(|| Scalar::list(Arc::clone(element_dtype), set, *nullability)) +} + fn compute_contains_scalar( list: &Scalar, needle: &Scalar, @@ -290,64 +360,13 @@ fn compute_list_contains( 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); + if !array.is::() { + return lists_contain_needles(array, 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. -/// -/// 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, -) -> 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| { - let comparison = Binary::try_new( - ConstantArray::new(element.clone(), len).into_array(), - values.clone(), - Operator::Eq, - )? - .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))?; - - 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 { - !elements.is_empty() - }; - 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) + let set = ListContainsSet::try_new(array, value.dtype(), options, ctx)? + .ok_or_else(|| vortex_err!("A non-null constant list of {} has a set", value.dtype()))?; + set.prepare(ctx)?.contains(value, ctx) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -439,6 +458,147 @@ fn list_contains_scalar( .into_array()) } +/// Neither side is constant: row `i` asks whether list `i` holds needle `i`. +/// +/// Each list's elements are gathered next to one copy of its row's needle, so a single equality +/// answers every comparison, and a fold over each list's contiguous range answers each row. +fn lists_contain_needles( + array: &ArrayRef, + values: &ArrayRef, + nullability: Nullability, + options: &ListContainsOptions, + ctx: &mut ExecutionCtx, +) -> VortexResult { + let list_array = array.clone().execute::(ctx)?; + let len = list_array.len(); + let list_valid = list_array + .validity()? + .execute_mask(len, ctx)? + .to_bit_buffer(); + + let offsets = list_array + .offsets() + .clone() + .execute::(ctx)?; + let offsets = offsets.reinterpret_cast(offsets.ptype().to_unsigned()); + let sizes = list_array.sizes().clone().execute::(ctx)?; + let sizes = sizes.reinterpret_cast(sizes.ptype().to_unsigned()); + let gathered = match_each_unsigned_integer_ptype!(offsets.ptype(), |O| { + match_each_unsigned_integer_ptype!(sizes.ptype(), |S| { + GatheredLists::new( + offsets.as_slice::(), + sizes.as_slice::(), + &list_valid, + ctx, + ) + }) + }); + + let elements = list_array + .elements() + .take(PrimitiveArray::new(gathered.element_idx, Validity::NonNullable).into_array())?; + let needles = + values.take(PrimitiveArray::new(gathered.row_idx, Validity::NonNullable).into_array())?; + let matches = Binary::try_new(elements, needles, Operator::Eq)? + .into_array() + .execute::(ctx)?; + let compared = matches + .validity()? + .execute_mask(matches.len(), ctx)? + .to_bit_buffer(); + + let starts = PrimitiveArray::new(gathered.starts, Validity::NonNullable); + let lens = PrimitiveArray::new(gathered.lens, Validity::NonNullable); + let any_true = process_matches::( + BoolArray::new(&matches.to_bit_buffer() & &compared, Validity::NonNullable), + len, + starts.clone(), + lens.clone(), + ctx, + ); + + // A null needle is null against a list with elements, but an empty list holds nothing to + // compare it to unless the function is strict. + let needle_valid = values.validity()?.execute_mask(len, ctx)?.to_bit_buffer(); + let mut decided = if options.sql_null_semantics { + needle_valid + } else { + &needle_valid | &!&non_empty_lists(&list_array, ctx)? + }; + // Under SQL null semantics a comparison with a null element leaves a non-match unknown. + if options.sql_null_semantics { + let any_null = process_matches::( + BoolArray::new(!&compared, Validity::NonNullable), + len, + starts, + lens, + ctx, + ); + decided = &decided & &(&any_true | &!&any_null); + } + + // A non-nullable result has non-null lists and needles, and no null element that could leave + // a row undecided. + let validity = match nullability { + Nullability::NonNullable => Validity::NonNullable, + Nullability::Nullable => list_array.validity()?.and(Validity::from(decided))?, + }; + Ok(BoolArray::new(any_true, validity).into_array()) +} + +/// The elements of every valid list laid out contiguously, each next to the row it belongs to. +struct GatheredLists { + /// The position of each gathered element in the list view's elements. + element_idx: Buffer, + /// The row, and so the needle, each gathered element is compared to. + row_idx: Buffer, + /// Where each row's run of gathered elements starts. + starts: Buffer, + /// How long each row's run is: the list's size, or zero for a null list. + lens: Buffer, +} + +impl GatheredLists { + fn new( + offsets: &[O], + sizes: &[S], + list_valid: &BitBuffer, + ctx: &mut ExecutionCtx, + ) -> Self { + let rows = sizes.len(); + let total: usize = (0..rows) + .filter(|&row| list_valid.value(row)) + .map(|row| sizes[row].as_()) + .sum(); + let allocator = ctx.allocator(); + let mut element_idx = BufferMut::::with_capacity_in(total, allocator.clone()); + let mut row_idx = BufferMut::::with_capacity_in(total, allocator.clone()); + let mut starts = BufferMut::::with_capacity_in(rows, allocator.clone()); + let mut lens = BufferMut::::with_capacity_in(rows, allocator.clone()); + + for row in 0..rows { + starts.push(element_idx.len() as u64); + // A null list's offset and size need not point at anything. + if !list_valid.value(row) { + lens.push(0); + continue; + } + let offset: usize = offsets[row].as_(); + let size: usize = sizes[row].as_(); + element_idx.extend((offset..offset + size).map(|idx| idx as u64)); + row_idx.extend(iter::repeat_n(row as u64, size)); + lens.push(size as u64); + } + + Self { + element_idx: element_idx.freeze(), + row_idx: row_idx.freeze(), + starts: starts.freeze(), + lens: lens.freeze(), + } + } +} + /// For each list, whether any set bit of `matches` falls in the list's element range. fn fold_lists( matches: BoolArray, @@ -593,6 +753,7 @@ mod tests { use std::sync::LazyLock; use itertools::Itertools; + use num_traits::PrimInt; use rstest::rstest; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; @@ -607,14 +768,19 @@ mod tests { use crate::VortexSessionExecute; use crate::array_session; use crate::arrays::BoolArray; + use crate::arrays::ChunkedArray; use crate::arrays::ConstantArray; + use crate::arrays::DictArray; use crate::arrays::ListArray; use crate::arrays::ListViewArray; use crate::arrays::PrimitiveArray; use crate::arrays::VarBinArray; + use crate::arrays::VarBinViewArray; use crate::assert_arrays_eq; use crate::dtype::DType; + use crate::dtype::NativePType; use crate::dtype::Nullability; + use crate::dtype::PType; use crate::dtype::PType::I32; use crate::dtype::StructFields; use crate::expr::Expression; @@ -630,9 +796,11 @@ mod tests { use crate::expr::or; use crate::expr::root; use crate::expr::stats::Stat; + use crate::scalar::PValue; use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::ListContains; use crate::scalar_fn::fns::list_contains::ListContainsOptions; + use crate::scalar_fn::fns::literal::Literal; use crate::stats::StatsSession; use crate::stats::stat as stat_expr; use crate::validity::Validity; @@ -1262,6 +1430,29 @@ mod tests { ) } + fn f64_set(values: Vec) -> Scalar { + Scalar::list( + Arc::new(DType::Primitive(PType::F64, Nullability::NonNullable)), + values.into_iter().map(Scalar::from).collect(), + Nullability::NonNullable, + ) + } + + fn utf8_set(values: Vec>) -> Scalar { + let element = DType::Utf8(Nullability::Nullable); + Scalar::list( + Arc::new(element.clone()), + values + .into_iter() + .map(|v| match v { + Some(v) => Scalar::utf8(v, Nullability::Nullable), + None => Scalar::null(element.clone()), + }) + .collect(), + Nullability::NonNullable, + ) + } + fn assert_result( result: VortexResult, expected: impl IntoIterator>, @@ -1271,6 +1462,33 @@ mod tests { Ok(()) } + #[test] + fn constant_set_of_floats_is_bitwise() -> VortexResult<()> { + // Membership has to agree with the compare kernel: `-0.0` and `0.0` are different + // members, and NaN is a member of a set that holds NaN. + let needles = PrimitiveArray::from_option_iter([ + Some(1.0f64), + Some(f64::NAN), + Some(0.0), + Some(-0.0), + Some(2.0), + None, + ]) + .into_array(); + let set = f64_set(vec![1.0, f64::NAN, -0.0]); + assert_result( + needles.apply(&list_contains(lit(set), root())), + [ + Some(true), + Some(true), + Some(false), + Some(true), + Some(false), + None, + ], + ) + } + #[test] fn constant_set_null_element_never_matches_by_default() -> VortexResult<()> { let needles = @@ -1372,6 +1590,233 @@ mod tests { ) } + /// Every row of `result` agrees with that row evaluated on its own, then matches `expected`. + fn assert_rows_agree(result: ArrayRef, expected: BoolArray) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let whole = result.clone().execute::(&mut ctx)?.into_array(); + for row in 0..result.len() { + assert_eq!( + result.execute_scalar(row, &mut ctx)?, + whole.execute_scalar(row, &mut ctx)?, + "row {row}" + ); + } + assert_arrays_eq!(whole, expected, &mut ctx); + Ok(()) + } + + #[rstest] + #[case::default( + ListContainsOptions::default(), + [Some(true), Some(false), None, Some(false), None, Some(true)] + )] + #[case::sql(SQL, [Some(true), None, None, None, None, Some(true)])] + fn list_column_against_needle_column( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 6], + ) -> VortexResult<()> { + // Lists `[1, 2]`, `[]`, null, `[3, null]`, `[4]`, `[5]` against needles 2, null, 1, 7, + // null, 5. + let lists = ListArray::try_new( + PrimitiveArray::from_option_iter([ + Some(1i32), + Some(2), + Some(3), + None, + Some(4), + Some(5), + ]) + .into_array(), + PrimitiveArray::from_iter(vec![0, 2, 2, 2, 4, 5, 6]).into_array(), + Validity::from_iter([true, true, false, true, true, true]), + )? + .into_array(); + let needles = + PrimitiveArray::from_option_iter([Some(2i32), None, Some(1), Some(7), None, Some(5)]) + .into_array(); + let result = ListContains::try_new_opts(lists, needles, options)?.into_array(); + assert_rows_agree(result, BoolArray::from_iter(expected)) + } + + #[test] + fn list_view_column_with_overlapping_views() -> VortexResult<()> { + // Views out of order and overlapping in `[1, 2, 3, 4]`: `[3, 4]` then `[1, 2, 3]`. + let lists = ListViewArray::try_new( + buffer![1i32, 2, 3, 4].into_array(), + buffer![2u32, 0].into_array(), + buffer![2u32, 3].into_array(), + Validity::NonNullable, + )? + .into_array(); + let needles = buffer![4i32, 4].into_array(); + let result = ListContains::try_new(lists, needles)?.into_array(); + assert_rows_agree(result, BoolArray::from_iter([true, false])) + } + + fn optimized_set(expr: &Expression) -> VortexResult { + let optimized = expr.optimize_recursive(&DType::Primitive(I32, Nullability::Nullable))?; + Ok(optimized + .child(0) + .as_opt::() + .vortex_expect("the set stays a literal") + .clone()) + } + + #[rstest] + // Off SQL null semantics a null element next to others decides nothing. + #[case::default(ListContainsOptions::default(), vec![Some(3), Some(1), None, Some(3), Some(2)], vec![Some(1), Some(2), Some(3)])] + // On them one null is kept, last, to leave a non-match unknown. + #[case::sql(SQL, vec![Some(3), None, Some(1), None, Some(3)], vec![Some(1), Some(3), None])] + // A list of nothing but nulls keeps one, since an empty list answers a null needle apart. + #[case::only_nulls(ListContainsOptions::default(), vec![None, None], vec![None])] + fn optimize_normalizes_the_set( + #[case] options: ListContainsOptions, + #[case] set: Vec>, + #[case] expected: Vec>, + ) -> VortexResult<()> { + let expr = list_contains_opts(lit(i32_set(set)), root(), options); + assert_eq!(optimized_set(&expr)?, i32_set(expected.clone())); + // A normalized set is left alone, so optimization reaches a fixed point. + let normalized = list_contains_opts(lit(i32_set(expected.clone())), root(), options); + assert_eq!(optimized_set(&normalized)?, i32_set(expected)); + Ok(()) + } + + #[rstest] + #[case::default(ListContainsOptions::default())] + #[case::sql(SQL)] + fn optimize_keeps_the_answer(#[case] options: ListContainsOptions) -> VortexResult<()> { + // Floats compare bitwise: `-0.0` and `0.0` stay distinct members, and NaN is one member. + let element = DType::Primitive(PType::F64, Nullability::Nullable); + let set = Scalar::list( + Arc::new(element.clone()), + [ + Some(2.0), + Some(f64::NAN), + None, + Some(-0.0), + Some(2.0), + Some(f64::NAN), + ] + .into_iter() + .map(|v| match v { + Some(v) => Scalar::primitive(v, Nullability::Nullable), + None => Scalar::null(element.clone()), + }) + .collect(), + Nullability::NonNullable, + ); + let needles = PrimitiveArray::from_option_iter([ + Some(2.0f64), + Some(0.0), + Some(-0.0), + Some(f64::NAN), + Some(5.0), + None, + ]) + .into_array(); + let expr = list_contains_opts(lit(set), root(), options); + let optimized = expr.optimize_recursive(needles.dtype())?; + assert_ne!(optimized, expr, "the set is normalized"); + + let mut ctx = array_session().create_execution_ctx(); + let expected = needles + .clone() + .apply(&expr)? + .execute::(&mut ctx)?; + assert_arrays_eq!(needles.apply(&optimized)?, expected, &mut ctx); + Ok(()) + } + + /// Probes `set` with each element, its neighbours and the type's extremes, against a naive + /// oracle. + fn assert_integer_membership(set: Vec) -> VortexResult<()> + where + T: NativePType + PrimInt + Into, + { + let mut needles = vec![T::min_value(), T::max_value()]; + for &element in &set { + needles.push(element); + needles.extend(element.checked_sub(&T::one())); + needles.extend(element.checked_add(&T::one())); + } + let expected = BoolArray::from_iter(needles.iter().map(|needle| set.contains(needle))); + let list = Scalar::list( + Arc::new(DType::Primitive(T::PTYPE, Nullability::NonNullable)), + set.into_iter() + .map(|v| Scalar::primitive(v, Nullability::NonNullable)) + .collect(), + Nullability::NonNullable, + ); + let mut ctx = array_session().create_execution_ctx(); + let result = PrimitiveArray::from_iter(needles) + .into_array() + .apply(&list_contains(lit(list), root()))?; + assert_arrays_eq!(result, expected, &mut ctx); + Ok(()) + } + + #[test] + fn integer_sets_through_every_probe() -> VortexResult<()> { + // Dense spans probe a bitmap, including spans straddling zero and whole types. + assert_integer_membership(vec![-3i32, -1, 0, 2])?; + assert_integer_membership(vec![i8::MIN, 0, i8::MAX])?; + assert_integer_membership(vec![u64::MAX, u64::MAX - 2])?; + assert_integer_membership(vec![i64::MIN, i64::MIN + 5])?; + // A sparse set probes a binary search when small and a hash set when large. + assert_integer_membership(vec![-(1i64 << 40), 1, 1 << 40])?; + assert_integer_membership((0..100u64).map(|v| v << 30).collect())?; + assert_integer_membership(vec![i64::MIN, i64::MAX]) + } + + #[rstest] + #[case::default(ListContainsOptions::default())] + #[case::sql(SQL)] + fn chunked_needles_agree_with_flat_needles( + #[case] options: ListContainsOptions, + ) -> VortexResult<()> { + // Chunks in different encodings, one of them empty, against a set with a null element. + let chunks = vec![ + PrimitiveArray::from_option_iter([Some(1i32), None, Some(4)]).into_array(), + PrimitiveArray::from_option_iter::([]).into_array(), + DictArray::try_new( + PrimitiveArray::from_option_iter([Some(0u32), None, Some(1)]).into_array(), + PrimitiveArray::from_option_iter([Some(2i32), Some(9)]).into_array(), + )? + .into_array(), + ]; + let dtype = chunks[0].dtype().clone(); + let chunked = ChunkedArray::try_new(chunks, dtype)?.into_array(); + let expr = list_contains_opts(lit(i32_set(vec![Some(2), None, Some(4)])), root(), options); + + let mut ctx = array_session().create_execution_ctx(); + let flat = + PrimitiveArray::from_option_iter([Some(1i32), None, Some(4), Some(2), None, Some(9)]) + .into_array() + .apply(&expr)? + .execute::(&mut ctx)?; + assert_rows_agree(chunked.apply(&expr)?, flat) + } + + #[test] + fn chunked_string_needles() -> VortexResult<()> { + let chunks = vec![ + VarBinViewArray::from_iter_nullable_str([Some("a"), None]).into_array(), + VarBinViewArray::from_iter_nullable_str([ + Some("a value longer than twelve bytes"), + Some("b"), + ]) + .into_array(), + ]; + let dtype = chunks[0].dtype().clone(); + let chunked = ChunkedArray::try_new(chunks, dtype)?.into_array(); + let set = utf8_set(vec![Some("a"), Some("a value longer than twelve bytes")]); + assert_rows_agree( + chunked.apply(&list_contains(lit(set), root()))?, + BoolArray::from_iter([Some(true), None, Some(true), Some(false)]), + ) + } + #[test] fn strict_only_under_sql_semantics() { let default = list_contains(lit(empty_i32_list()), root()); @@ -1399,6 +1844,129 @@ mod tests { assert_result(needles.apply(&in_list(root(), lit(set))), [None, None]) } + #[test] + fn constant_set_row_probe_keeps_the_declared_nullability() -> VortexResult<()> { + // A nullable element dtype decides nothing off SQL null semantics, so non-null needles give + // a non-nullable result, though each comparison is against a nullable element. + let needles = BoolArray::from_iter([true, false]).into_array(); + let element = DType::Bool(Nullability::Nullable); + let set = Scalar::list( + Arc::new(element.clone()), + vec![ + Scalar::bool(true, Nullability::Nullable), + Scalar::null(element), + ], + Nullability::NonNullable, + ); + let result = needles.apply(&list_contains(lit(set), root()))?; + assert_eq!(result.dtype(), &DType::Bool(Nullability::NonNullable)); + assert_rows_agree(result, BoolArray::from_iter([true, false])) + } + + #[test] + fn constant_set_row_probe_of_only_nulls() -> VortexResult<()> { + // The row probe must also preserve the distinction between an empty set and null elements. + let needles = BoolArray::from_iter([Some(true), None]).into_array(); + let element = DType::Bool(Nullability::Nullable); + let set = Scalar::list( + Arc::new(element.clone()), + vec![Scalar::null(element.clone()), Scalar::null(element)], + Nullability::NonNullable, + ); + 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 constant_set_of_strings_reaches_out_of_line_values() -> VortexResult<()> { + // Values longer than 12 bytes live outside the view; both kinds must be looked up. + let long = "a value longer than twelve bytes"; + let needles = VarBinViewArray::from_iter_nullable_str([ + Some("a"), + Some(long), + Some("a value longer than twelve byteX"), + Some("b"), + None, + ]) + .into_array(); + let set = utf8_set(vec![Some("a"), Some(long)]); + assert_result( + needles.clone().apply(&list_contains(lit(set), root())), + [Some(true), Some(true), Some(false), Some(false), None], + )?; + let set_with_null = utf8_set(vec![Some("a"), None]); + assert_result( + needles.apply(&in_list(root(), lit(set_with_null))), + [Some(true), None, None, None, None], + ) + } + + #[rstest] + // Non-strict, so a null code keeps the dictionary rewrite from pushing the function into the + // values: the needle is executed into its canonical encoding and probed against the set. + #[case::executed(ListContainsOptions::default(), [Some(false), Some(true), None, Some(true)])] + // Strict, so the rewrite probes the dictionary's two values instead of its four rows. + #[case::pushed_into_values(SQL, [Some(false), Some(true), None, Some(true)])] + fn constant_set_probes_a_dictionary_needle( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 4], + ) -> VortexResult<()> { + let needles = DictArray::try_new( + PrimitiveArray::from_option_iter([Some(0u32), Some(1), None, Some(1)]).into_array(), + PrimitiveArray::from_iter([1i32, 2]).into_array(), + )? + .into_array(); + let set = i32_set(vec![Some(2), Some(4)]); + assert_rows_agree( + needles.apply(&list_contains_opts(lit(set), root(), options))?, + BoolArray::from_iter(expected), + ) + } + + #[test] + fn constant_set_of_strings_probes_an_encoded_needle() -> VortexResult<()> { + let needles = VarBinArray::from_iter( + [Some("a"), Some("b"), None], + DType::Utf8(Nullability::Nullable), + ) + .into_array(); + let set = utf8_set(vec![Some("a"), Some("c")]); + assert_result( + needles.apply(&list_contains(lit(set), root())), + [Some(true), Some(false), None], + ) + } + + #[test] + fn constant_set_row_probe_honours_both_semantics() -> VortexResult<()> { + // Boolean sets use the general sorted row probe. + let needles = BoolArray::from_iter([Some(true), Some(false), None]).into_array(); + let element = DType::Bool(Nullability::Nullable); + let set = Scalar::list( + Arc::new(element.clone()), + vec![ + Scalar::bool(true, Nullability::Nullable), + Scalar::null(element), + ], + Nullability::NonNullable, + ); + assert_result( + needles + .clone() + .apply(&list_contains(lit(set.clone()), root())), + [Some(true), Some(false), None], + )?; + assert_result( + needles.apply(&in_list(root(), lit(set))), + [Some(true), None, None], + ) + } + #[test] fn list_array_null_elements_under_sql_semantics() -> VortexResult<()> { // Lists `[1, null]`, `[2]`, `[3]`, `[]` against a constant needle. diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs new file mode 100644 index 00000000000..e87690d1c3f --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -0,0 +1,487 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! Probing a constant set: the probe structure is built from a [`ListContainsSet`] for canonical +//! evaluation. Chunked expressions dispatch per chunk so each encoding can specialize membership. + +use std::hash::BuildHasher; + +use num_traits::ToPrimitive; +use num_traits::WrappingSub; +use vortex_buffer::BitBuffer; +use vortex_buffer::BitBufferMut; +use vortex_buffer::Buffer; +use vortex_buffer::BufferAllocatorRef; +use vortex_buffer::BufferMut; +use vortex_error::VortexExpect; +use vortex_error::VortexResult; +use vortex_utils::aliases::hash_map::HashTable; +use vortex_utils::aliases::hash_map::HashTableEntry; +use vortex_utils::aliases::hash_map::RandomState; + +use crate::ArrayRef; +use crate::ExecutionCtx; +use crate::IntoArray; +use crate::RecursiveCanonical; +use crate::arrays::DecimalArray; +use crate::arrays::PrimitiveArray; +use crate::arrays::VarBinViewArray; +use crate::arrays::decimal::DecimalArrayExt; +use crate::arrays::decimal::widened_buffer; +use crate::arrays::primitive::PrimitiveArrayExt; +use crate::arrays::varbinview::BinaryView; +use crate::dtype::DType; +use crate::dtype::IntegerPType; +use crate::dtype::PType; +use crate::dtype::i256; +use crate::match_each_decimal_value_type; +use crate::match_each_integer_ptype; +use crate::scalar::DecimalValue; +use crate::scalar_fn::fns::binary::build_row_comparator; +use crate::scalar_fn::fns::binary::collect_bits; +use crate::scalar_fn::fns::list_contains::ListContainsSet; +use crate::validity::Validity; + +/// A set whose span of values needs at most this many bits per element is probed through a bitmap +/// over the span, bounding the bitmap to a few words per element. +const BITMAP_BITS_PER_ELEMENT: u128 = 64; +/// A span this narrow is probed through a bitmap whatever the size of the set. +const BITMAP_MIN_BITS: u128 = 1 << 12; + +/// A [`ListContainsSet`] with the structure that probes it, prepared once and then run over any +/// number of needle arrays of the set's dtype. +/// +/// Every needle is probed against a bitmap, a sorted set or a hash table. Decimals use their +/// unscaled integers with the same bitmap and sorted-value strategy as primitives. Nested +/// values use sorted row indices with the same comparator as equality. No probe constructs +/// per-element expressions or materializes scalars in its loop. +pub struct PreparedSet { + set: ListContainsSet, + probe: Probe, +} + +enum Probe { + /// Integers, or floats by their bit patterns, spanning a dense range: one bit per value of the + /// span above `min_offset`, the smallest element as a `usize`. + Bitmap { + min_offset: usize, + span: usize, + bitmap: BitBuffer, + }, + /// Integers, or floats by their bit patterns, sorted without duplicates. + Sorted(PrimitiveArray), + /// Dense unscaled decimal values, including 128- and 256-bit storage. + DecimalBitmap { + min: DecimalValue, + bitmap: BitBuffer, + }, + /// Sorted, distinct unscaled decimal values. + DecimalSorted(DecimalArray), + /// UTF-8 or binary elements, found through a table of their indices hashed by their bytes, so + /// that no element is copied. + Bytes { + elements: VarBinViewArray, + hasher: RandomState, + table: HashTable, + }, + /// Recursively canonical elements, indexed in sorted order with duplicates removed. + Rows { + elements: ArrayRef, + indices: Vec, + }, +} + +impl ListContainsSet { + /// Builds the structure that probes this set. + pub fn prepare(self, ctx: &mut ExecutionCtx) -> VortexResult { + let probe = match self.elements().dtype() { + DType::Primitive(ptype, _) => { + let ptype = bit_pattern_ptype(*ptype); + let elements = self + .elements() + .clone() + .execute::(ctx)? + .reinterpret_cast(ptype); + match_each_integer_ptype!(ptype, |T| { + integer_probe::(elements, ctx.allocator()) + }) + } + DType::Decimal(..) => { + let elements = self.elements().clone().execute::(ctx)?; + match_each_decimal_value_type!(elements.values_type(), |T| { + let values = elements.buffer::(); + if let Some((min, bitmap)) = integer_bitmap(&values, ctx.allocator()) { + Probe::DecimalBitmap { + min: min.into(), + bitmap, + } + } else { + Probe::DecimalSorted(DecimalArray::new( + sorted_values(values), + elements.decimal_dtype(), + Validity::NonNullable, + )) + } + }) + } + DType::Utf8(_) | DType::Binary(_) => { + bytes_probe(self.elements().clone().execute::(ctx)?) + } + _ => { + let elements = self + .elements() + .clone() + .execute::(ctx)? + .0 + .into_array(); + let mut indices: Vec = (0..elements.len()).collect(); + if !indices.is_empty() { + let compare = build_row_comparator(&elements, &elements, ctx)?; + if !indices.is_sorted_by(|&lhs, &rhs| compare(lhs, rhs).is_le()) { + indices.sort_unstable_by(|&lhs, &rhs| compare(lhs, rhs)); + } + indices.dedup_by(|lhs, rhs| compare(*lhs, *rhs).is_eq()); + } + Probe::Rows { elements, indices } + } + }; + Ok(PreparedSet { set: self, probe }) + } +} + +impl PreparedSet { + /// The dtype of every result, as the expression declares it. + pub fn dtype(&self) -> DType { + DType::Bool(self.set.nullability) + } + + /// Whether each of `needles`, which have the set's dtype, is a member of the set. + /// + /// The result has the nullability the expression declares, whatever the needles' encoding. + pub fn contains(&self, needles: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { + match &self.probe { + Probe::Bitmap { .. } | Probe::Sorted(_) => self.contains_primitive(needles, ctx), + Probe::DecimalBitmap { .. } | Probe::DecimalSorted(_) => { + self.contains_decimal(needles, ctx) + } + Probe::Bytes { + elements, + hasher, + table, + } => self.contains_bytes(elements, hasher, table, needles, ctx), + Probe::Rows { elements, indices } => { + self.contains_rows(elements, indices, needles, ctx) + } + } + } + + /// A float is a member exactly when the compare kernel would call it equal to an element, which + /// is when their bit patterns match — distinguishing `-0.0` from `0.0` and one NaN payload from + /// another — so floats are probed by their bits, as integers. + fn contains_primitive( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let primitive = needles.clone().execute::(ctx)?; + let ptype = bit_pattern_ptype(primitive.ptype()); + let values = primitive.reinterpret_cast(ptype); + let bits = match_each_integer_ptype!(ptype, |T| { + self.probe + .integer_bits(values.as_slice::(), ctx.allocator()) + }); + self.set.finish(bits, primitive.validity()?) + } + + fn contains_decimal( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let needles = needles.clone().execute::(ctx)?; + // Logical precision and scale agree, but physical widths can differ. Widen once to + // the common width so out-of-range needles cannot truncate into false matches. + let bits = match &self.probe { + Probe::DecimalBitmap { min, bitmap } => { + let common = min.decimal_type().max(needles.values_type()); + match_each_decimal_value_type!(common, |T| { + let min = min.cast::().vortex_expect("lossless decimal widening"); + let values = widened_buffer::(&needles); + collect_bits( + &values, + |value| { + value.offset_from(min).is_some_and(|offset| { + offset < bitmap.len() && bitmap.value(offset) + }) + }, + ctx.allocator(), + ) + }) + } + Probe::DecimalSorted(sorted) => { + let common = sorted.values_type().max(needles.values_type()); + match_each_decimal_value_type!(common, |T| { + let sorted = widened_buffer::(sorted); + let values = widened_buffer::(&needles); + sorted_bits(&sorted, &values, ctx.allocator()) + }) + } + _ => unreachable!("decimal needles meet a decimal probe"), + }; + self.set.finish(bits, needles.validity()?) + } + + fn contains_bytes( + &self, + elements: &VarBinViewArray, + hasher: &RandomState, + table: &HashTable, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let element_views = elements.views(); + let element_buffers = data_buffers(elements); + let array = needles.clone().execute::(ctx)?; + let buffers = data_buffers(&array); + let bits = collect_bits( + array.views(), + |view: BinaryView| { + let value = view_bytes(&view, &buffers); + table + .find(hasher.hash_one(value), |&idx| { + view_bytes(&element_views[idx as usize], &element_buffers) == value + }) + .is_some() + }, + ctx.allocator(), + ); + self.set.finish(bits, array.validity()?) + } + + fn contains_rows( + &self, + elements: &ArrayRef, + indices: &[usize], + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + if indices.is_empty() { + return self.set.finish( + BitBuffer::full_in(false, needles.len(), ctx.allocator().clone()), + needles.validity()?, + ); + } + let needles = needles + .clone() + .execute::(ctx)? + .0 + .into_array(); + // The comparator materializes child buffers and validity once, including decimal + // widening and nested offsets. Each search then reads those buffers directly. + let compare = build_row_comparator(elements, &needles, ctx)?; + let bits = BitBuffer::collect_bool_in( + needles.len(), + |row| { + indices + .binary_search_by(|&element| compare(element, row)) + .is_ok() + }, + ctx.allocator().clone(), + ); + self.set.finish(bits, needles.validity()?) + } +} + +impl Probe { + /// One bit per needle, set when the needle is an element. + fn integer_bits( + &self, + needles: &[T], + allocator: &BufferAllocatorRef, + ) -> BitBuffer { + match self { + Self::Bitmap { + min_offset, + span, + bitmap, + } => collect_bits( + needles, + // A needle below the smallest element wraps past `span`, so one comparison checks + // both bounds. + |needle| { + let offset = needle.as_().wrapping_sub(*min_offset); + offset <= *span && bitmap.value(offset) + }, + allocator, + ), + Self::Sorted(sorted) => { + let sorted = sorted.as_slice::(); + sorted_bits(sorted, needles, allocator) + } + Self::DecimalBitmap { .. } + | Self::DecimalSorted(_) + | Self::Bytes { .. } + | Self::Rows { .. } => { + unreachable!("integer needles meet an integer probe") + } + } + } +} + +/// A table of the elements' indices, hashed by their bytes, holding each distinct value once. +fn bytes_probe(elements: VarBinViewArray) -> Probe { + let hasher = RandomState::default(); + let mut table = HashTable::with_capacity(elements.len()); + { + let views = elements.views(); + let buffers = data_buffers(&elements); + let bytes = |idx: u32| view_bytes(&views[idx as usize], &buffers); + for (idx, view) in views.iter().enumerate() { + let value = view_bytes(view, &buffers); + if let HashTableEntry::Vacant(vacant) = table.entry( + hasher.hash_one(value), + |&other| bytes(other) == value, + |&other| hasher.hash_one(bytes(other)), + ) { + vacant.insert( + u32::try_from(idx).vortex_expect("a list holds fewer than 2^32 elements"), + ); + } + } + } + Probe::Bytes { + elements, + hasher, + table, + } +} + +/// The host slices of an array's data buffers, indexed by a view's buffer index. +fn data_buffers(array: &VarBinViewArray) -> Vec<&[u8]> { + (0..array.data_buffers().len()) + .map(|idx| array.buffer(idx).as_slice()) + .collect() +} + +/// The bytes a view points at, inlined in the view itself or out of line in one of `buffers`. +fn view_bytes<'a>(view: &'a BinaryView, buffers: &[&'a [u8]]) -> &'a [u8] { + if view.is_inlined() { + view.as_inlined().value() + } else { + let reference = view.as_view(); + &buffers[reference.buffer_index as usize][reference.as_range()] + } +} + +/// The integer type with a float's bit pattern, or the type itself for an integer. +fn bit_pattern_ptype(ptype: PType) -> PType { + match ptype { + PType::F16 => PType::U16, + PType::F32 => PType::U32, + PType::F64 => PType::U64, + _ => ptype, + } +} + +/// A bitmap over the elements' span when the span is dense, and a sorted slice otherwise. +/// +/// A hash set and, for a handful of elements, a linear scan both lost to the binary search at every +/// set size measured by the `list_contains_set` benchmark, up to 16 384 elements. +fn integer_probe( + elements: PrimitiveArray, + allocator: &BufferAllocatorRef, +) -> Probe { + // The primitive hot loop computes offsets in usize, which must hold every value of T. + if size_of::() <= size_of::() + && let Some((min, bitmap)) = integer_bitmap(elements.as_slice::(), allocator) + { + return Probe::Bitmap { + min_offset: min.as_(), + span: bitmap.len() - 1, + bitmap, + }; + } + Probe::Sorted(PrimitiveArray::new( + sorted_values(elements.into_buffer::()), + Validity::NonNullable, + )) +} + +/// An unsigned modular distance, rejecting offsets too wide to address a bitmap. +/// Keeping the subtraction at the physical width also handles signed ranges spanning zero. +trait SetInteger: Copy + Ord { + fn offset_from(self, min: Self) -> Option; +} + +macro_rules! impl_set_integer { + ($($signed:ty => $unsigned:ty),* $(,)?) => { + $(impl SetInteger for $signed { + fn offset_from(self, min: Self) -> Option { + usize::try_from((self as $unsigned).wrapping_sub(min as $unsigned)).ok() + } + })* + }; +} + +impl_set_integer!( + u8 => u8, u16 => u16, u32 => u32, u64 => u64, + i8 => u8, i16 => u16, i32 => u32, i64 => u64, i128 => u128, +); + +impl SetInteger for i256 { + fn offset_from(self, min: Self) -> Option { + self.wrapping_sub(&min).to_usize() + } +} + +fn integer_bitmap( + values: &[T], + allocator: &BufferAllocatorRef, +) -> Option<(T, BitBuffer)> { + let min = *values.iter().min()?; + let max = *values.iter().max()?; + let span = max.offset_from(min)?; + if span == usize::MAX + || span as u128 >= (values.len() as u128 * BITMAP_BITS_PER_ELEMENT).max(BITMAP_MIN_BITS) + { + return None; + } + let mut bitmap = BitBufferMut::from_buffer( + BufferMut::zeroed_in((span + 1).div_ceil(8), allocator.clone()), + 0, + span + 1, + ); + for &value in values { + bitmap.set( + value + .offset_from(min) + .vortex_expect("value within bitmap span"), + ); + } + Some((min, bitmap.freeze())) +} + +fn sorted_values(values: Buffer) -> Buffer { + if values.is_sorted_by(|a, b| a < b) { + return values; + } + let mut sorted = values.to_vec(); + sorted.sort_unstable(); + sorted.dedup(); + Buffer::from(sorted) +} + +fn sorted_bits( + sorted: &[T], + needles: &[T], + allocator: &BufferAllocatorRef, +) -> BitBuffer { + collect_bits( + needles, + |needle| sorted.binary_search(&needle).is_ok(), + allocator, + ) +} + +#[cfg(test)] +mod tests; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs new file mode 100644 index 00000000000..68dc02194ed --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -0,0 +1,226 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use rstest::rstest; +use vortex_buffer::BitBuffer; +use vortex_buffer::buffer; +use vortex_error::VortexResult; + +use super::Probe; +use crate::ArrayRef; +use crate::IntoArray; +use crate::array_session; +use crate::arrays::Bool; +use crate::arrays::BoolArray; +use crate::arrays::ConstantArray; +use crate::arrays::DecimalArray; +use crate::arrays::FixedSizeListArray; +use crate::arrays::ListArray; +use crate::arrays::PrimitiveArray; +use crate::arrays::StructArray; +use crate::assert_arrays_eq; +use crate::builders::builder_with_capacity_in; +use crate::dtype::DType; +use crate::dtype::DecimalDType; +use crate::dtype::DecimalType; +use crate::dtype::Nullability; +use crate::dtype::PType; +use crate::dtype::i256; +use crate::match_each_decimal_value_type; +use crate::scalar::DecimalValue; +use crate::scalar::Scalar; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::scalar_fn::fns::list_contains::ListContainsSet; +use crate::validity::Validity; + +fn nested_needles() -> ArrayRef { + ListArray::try_new( + PrimitiveArray::from_option_iter([Some(1i32), None, Some(2), None, Some(9)]) + .into_array(), + buffer![0u32, 2, 2, 4, 5].into_array(), + Validity::from(BitBuffer::from_iter([true, false, true, true])), + ) + .unwrap() + .into_array() +} + +fn map_needles() -> ArrayRef { + let ctx = array_session().create_execution_ctx(); + let dtype = DType::map( + DType::Primitive(PType::I32, Nullability::NonNullable), + DType::Utf8(Nullability::Nullable), + false, + Nullability::Nullable, + ) + .unwrap(); + let mut builder = builder_with_capacity_in(&dtype, 4, ctx.allocator()); + for key in [Some(1i32), None, Some(2), Some(9)] { + let scalar = match key { + Some(key) => Scalar::map( + dtype.clone(), + [(key.into(), Scalar::null(DType::Utf8(Nullability::Nullable)))], + ), + None => Scalar::null(dtype.clone()), + }; + builder.append_scalar(&scalar).unwrap(); + } + builder.finish() +} + +#[rstest] +#[case::list(nested_needles())] +#[case::map(map_needles())] +#[case::struct_of_lists( + StructArray::from_fields(&[("list", nested_needles())]).unwrap().into_array() +)] +#[case::fixed_size_list(FixedSizeListArray::new( + PrimitiveArray::from_option_iter([ + Some(1i32), None, Some(2), None, Some(9), None, Some(8), None, + ]).into_array(), + 2, Validity::NonNullable, 4, +).into_array())] +fn test_row_set_returns_membership_bits( + #[case] needles: ArrayRef, + #[values(false, true)] sql_null_semantics: bool, +) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let dtype = needles.dtype().as_nullable(); + // Unsorted, repeated members and a top-level null exercise normalization and SQL's + // unknown non-match independently of any nulls nested inside the members. + let members = [2, 0, 2] + .map(|idx| needles.execute_scalar(idx, &mut ctx)?.cast(&dtype)) + .into_iter() + .collect::>>()?; + let mut elements = members.clone(); + elements.push(Scalar::null(dtype.clone())); + let list = ConstantArray::new( + Scalar::list(dtype, elements, Nullability::NonNullable), + needles.len(), + ) + .into_array(); + let options = ListContainsOptions { sql_null_semantics }; + let set = ListContainsSet::try_new(&list, needles.dtype(), &options, &mut ctx)? + .unwrap() + .prepare(&mut ctx)?; + let result = set.contains(&needles, &mut ctx)?; + // The old fallback returned a lazy OR tree here. + assert!(result.is::()); + let expected = (0..needles.len()) + .map(|idx| { + let value = needles.execute_scalar(idx, &mut ctx)?; + Ok(if value.is_null() { + None + } else if members.contains(&value) { + Some(true) + } else if sql_null_semantics { + None + } else { + Some(false) + }) + }) + .collect::>>()?; + assert_arrays_eq!(result, BoolArray::from_iter(expected), &mut ctx); + Ok(()) +} + +#[rstest] +fn test_decimal_bitmap_across_storage_widths( + #[values(2, 4, 9, 18, 38, 76)] precision: u8, + #[values( + DecimalType::I8, DecimalType::I16, DecimalType::I32, + DecimalType::I64, DecimalType::I128, DecimalType::I256, + )] + needle_width: DecimalType, + #[values(false, true)] sql_null_semantics: bool, +) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal = DecimalDType::new(precision, 1); + let dtype = DType::Decimal(decimal, Nullability::Nullable); + let mut elements: Vec<_> = [1i8, -2, 1] + .map(|value| Scalar::decimal(value.into(), decimal, Nullability::Nullable)) + .into(); + elements.push(Scalar::null(dtype.clone())); + let list = ConstantArray::new( + Scalar::list(dtype, elements, Nullability::NonNullable), + 4, + ) + .into_array(); + let needles = match_each_decimal_value_type!(needle_width, |T| { + DecimalArray::from_option_iter::( + [Some(-2i8), Some(0), Some(1), None].map(|value| { + value.map(|value| DecimalValue::from(value).cast::().unwrap()) + }), + decimal, + ) + .into_array() + }); + let options = ListContainsOptions { sql_null_semantics }; + let set = ListContainsSet::try_new(&list, needles.dtype(), &options, &mut ctx)? + .unwrap() + .prepare(&mut ctx)?; + assert!(matches!(&set.probe, Probe::DecimalBitmap { .. })); + let non_match = (!sql_null_semantics).then_some(false); + assert_arrays_eq!( + set.contains(&needles, &mut ctx)?, + BoolArray::from_iter([Some(true), non_match, Some(true), None]), + &mut ctx, + ); + Ok(()) +} + +#[rstest] +#[case::dense_i128(38, i256::from_i128(1i128 << 100), true)] +#[case::sparse_i128(38, i256::from_i128(1i128 << 100), false)] +#[case::dense_i256(76, i256::from_parts(0, 1i128 << 72), true)] +#[case::sparse_i256(76, i256::from_parts(0, 1i128 << 72), false)] +fn test_decimal_wide_values( + #[case] precision: u8, + #[case] base: i256, + #[case] dense: bool, +) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let decimal = DecimalDType::new(precision, 2); + let second = if dense { base + i256::ONE } else { -base }; + let list = ConstantArray::new( + Scalar::list( + DType::Decimal(decimal, Nullability::NonNullable), + [second, base, second] + .map(|v| Scalar::decimal(v.into(), decimal, Nullability::NonNullable)) + .into(), + Nullability::NonNullable, + ), + 6, + ) + .into_array(); + // The last non-null value shares the low 64 bits of a member but must not match it. + let needles = DecimalArray::from_option_iter::( + [ + Some(base), + Some(second), + Some(base - i256::ONE), + Some(i256::ZERO), + Some(base + i256::from_i128(1i128 << 64)), + None, + ], + decimal, + ) + .into_array(); + let set = ListContainsSet::try_new( + &list, + needles.dtype(), + &ListContainsOptions::default(), + &mut ctx, + )? + .unwrap() + .prepare(&mut ctx)?; + assert_eq!(matches!(&set.probe, Probe::DecimalBitmap { .. }), dense); + assert_eq!(matches!(&set.probe, Probe::DecimalSorted(_)), !dense); + assert_arrays_eq!( + set.contains(&needles, &mut ctx)?, + BoolArray::from_iter([ + Some(true), Some(true), Some(false), Some(false), Some(false), None, + ]), + &mut ctx, + ); + Ok(()) +} From 84e5d8c264a3e2251445c64be36445645249fdf9 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 11:01:31 -0400 Subject: [PATCH 2/9] Fix prepared-set test compilation Import the execution-context trait, remove trailing macro commas, and construct the struct fixture outside the generated fallible test wrapper. Signed-off-by: Robert Kruszewski --- .../scalar_fn/fns/list_contains/prepared/tests.rs | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs index 68dc02194ed..c0faa84440a 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -9,6 +9,7 @@ use vortex_error::VortexResult; use super::Probe; use crate::ArrayRef; use crate::IntoArray; +use crate::VortexSessionExecute; use crate::array_session; use crate::arrays::Bool; use crate::arrays::BoolArray; @@ -67,12 +68,16 @@ fn map_needles() -> ArrayRef { builder.finish() } +fn struct_needles() -> ArrayRef { + StructArray::from_fields(&[("list", nested_needles())]) + .unwrap() + .into_array() +} + #[rstest] #[case::list(nested_needles())] #[case::map(map_needles())] -#[case::struct_of_lists( - StructArray::from_fields(&[("list", nested_needles())]).unwrap().into_array() -)] +#[case::struct_of_lists(struct_needles())] #[case::fixed_size_list(FixedSizeListArray::new( PrimitiveArray::from_option_iter([ Some(1i32), None, Some(2), None, Some(9), None, Some(8), None, @@ -163,7 +168,7 @@ fn test_decimal_bitmap_across_storage_widths( assert_arrays_eq!( set.contains(&needles, &mut ctx)?, BoolArray::from_iter([Some(true), non_match, Some(true), None]), - &mut ctx, + &mut ctx ); Ok(()) } @@ -220,7 +225,7 @@ fn test_decimal_wide_values( BoolArray::from_iter([ Some(true), Some(true), Some(false), Some(false), Some(false), None, ]), - &mut ctx, + &mut ctx ); Ok(()) } From ea29682e04e9d47ddd3ef7da6067ea79cfe7d51d Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 12:27:52 -0400 Subject: [PATCH 3/9] Format prepared membership probes and tests Signed-off-by: Robert Kruszewski --- .../src/scalar_fn/fns/list_contains/kernel.rs | 116 +----- .../src/scalar_fn/fns/list_contains/mod.rs | 271 ++++++------ .../fns/list_contains/prepared/array.rs | 387 ++++++++++++++++++ .../fns/list_contains/prepared/mod.rs | 103 ++--- .../fns/list_contains/prepared/tests.rs | 175 +++++--- 5 files changed, 693 insertions(+), 359 deletions(-) create mode 100644 vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs 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 cab134e6f5f..b4ba2ee520b 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -1,29 +1,21 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -use vortex_buffer::BitBuffer; use vortex_error::VortexExpect; use vortex_error::VortexResult; use crate::ArrayRef; use crate::ExecutionCtx; -use crate::IntoArray; use crate::array::ArrayView; use crate::array::VTable; -use crate::arrays::BoolArray; -use crate::arrays::Constant; use crate::arrays::ScalarFn; -use crate::arrays::constant::list_scalar_elements; use crate::arrays::scalar_fn::ExactScalarFn; use crate::arrays::scalar_fn::ScalarFnArrayExt; use crate::arrays::scalar_fn::ScalarFnArrayView; -use crate::dtype::DType; -use crate::dtype::Nullability; 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; -use crate::validity::Validity; /// Check list-contains without reading buffers (metadata-only). /// @@ -45,6 +37,14 @@ pub trait ListContainsElementReduce: VTable { /// Like [`ListContainsElementReduce`], this dispatches on the **element** (needle) child at /// index 1. Unlike the reduce variant, implementations may read and execute on buffers via /// the provided [`ExecutionCtx`]. +/// +/// For a needle that is not canonical, execution prepares a constant list into a +/// [`PreparedSetArray`] and runs the kernels again. Thus a kernel can get the prepared set with +/// `list.as_opt::()`, and probe its own values with [`PreparedSetData::contains`]. +/// For example, a dictionary probes only its values. +/// +/// [`PreparedSetArray`]: crate::scalar_fn::fns::list_contains::PreparedSetArray +/// [`PreparedSetData::contains`]: crate::scalar_fn::fns::list_contains::PreparedSetData::contains pub trait ListContainsElementKernel: VTable { fn list_contains( list: &ArrayRef, @@ -54,106 +54,6 @@ pub trait ListContainsElementKernel: VTable { ) -> VortexResult>; } -/// The constant haystack of a `list_contains`, prepared for one pass over a column of needles. -/// -/// This is the `IN` shape, `list_contains(lit([...]), column)`: the list's elements are written -/// once into the array [`elements`](Self::elements), a kernel builds a probe structure from them and -/// tests every needle against it, then hands the membership bits back to [`finish`](Self::finish). -/// Shared by the canonical implementations and the element kernels, so that the result nullability -/// and the SQL null semantics have one definition. -/// -/// A null element is dropped from the elements — it never equals a needle — and only decides -/// whether a non-match is [unknown](Self::non_match_is_unknown). An empty list is an empty set, -/// whose answer for a null needle [`finish`](Self::finish) settles by the options. -/// -/// Preparing the set allocates and executes, so this serves [`ListContainsElementKernel`]. A -/// [`ListContainsElementReduce`] rule, which may not read buffers, works from the list's scalar. -pub struct ListContainsSet { - elements: ArrayRef, - pub(super) nullability: Nullability, - non_match_is_unknown: bool, - /// Whether every needle, a null one included, is absent: an empty list off SQL null semantics. - matches_nothing: bool, -} - -impl ListContainsSet { - /// Prepares the constant list `list` to be probed by needles of dtype `needle_dtype`. - /// - /// Returns `None` when there is no set to probe: a list that is not a constant or is null, or a - /// needle whose dtype does not match the list's elements. - pub fn try_new( - list: &ArrayRef, - needle_dtype: &DType, - options: &ListContainsOptions, - ctx: &mut ExecutionCtx, - ) -> VortexResult> { - let DType::List(element_dtype, _) = list.dtype() else { - return Ok(None); - }; - // A needle of a different dtype is an invalid expression, which the scalar function - // reports as an error. - if !element_dtype.eq_ignore_nullability(needle_dtype) { - return Ok(None); - } - // Every needle is probed against the same set, so the haystack has to be one constant - // list, and a null one holds no elements at all. - let Some(constant) = list.as_opt::() else { - return Ok(None); - }; - let list_scalar = constant.scalar(); - if list_scalar.is_null() { - return Ok(None); - } - - let elements = list_scalar_elements(&list_scalar.as_list(), ctx.allocator()); - let size = elements.len(); - - let valid = elements.validity()?.execute_mask(size, ctx)?; - let non_match_is_unknown = options.sql_null_semantics && !valid.all_true(); - // A null element never equals a needle, so a probe is better off without it. - let elements = if valid.all_true() { - elements - } else { - elements.filter(valid)? - }; - - Ok(Some(Self { - elements, - nullability: options.result_nullability(list.dtype(), needle_dtype), - non_match_is_unknown, - matches_nothing: size == 0 && !options.sql_null_semantics, - })) - } - - /// The list's non-null elements, one row each, in the dtype of the needles. - pub fn elements(&self) -> &ArrayRef { - &self.elements - } - - /// Whether a needle matching no element answers `null` rather than `false`. - /// - /// A null element under [`ListContainsOptions::sql_null_semantics`] makes the comparison - /// unknown, the three-valued semantics of SQL `IN`. - pub fn non_match_is_unknown(&self) -> bool { - self.non_match_is_unknown - } - - /// Assembles the result from one membership bit per needle and the needles' validity. - /// - /// A null needle stays null, unless the list is empty off SQL null semantics, which answers - /// `false` for every needle. When a non-match is unknown, only the rows that matched stay valid. - pub fn finish(&self, bits: BitBuffer, needle_validity: Validity) -> VortexResult { - let validity = if self.matches_nothing { - Validity::NonNullable - } else if self.non_match_is_unknown { - needle_validity.and(Validity::from(bits.clone()))? - } else { - needle_validity - }; - Ok(BoolArray::new(bits, validity.union_nullability(self.nullability)).into_array()) - } -} - /// Adaptor that wraps a [`ListContainsElementReduce`] impl as an [`ArrayParentReduceRule`]. #[derive(Default, Debug)] pub struct ListContainsElementReduceAdaptor(pub V); 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 e1712d835e9..d23906a0b06 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -4,18 +4,18 @@ mod kernel; mod prepared; -use std::cmp::Ordering; use std::fmt::Display; use std::fmt::Formatter; use std::hash::Hash; use std::iter; use std::ops::BitOr; -use std::sync::Arc; use arrow_buffer::bit_iterator::BitIndexIterator; pub use kernel::*; use num_traits::Zero; pub use prepared::PreparedSet; +pub use prepared::PreparedSetArray; +pub use prepared::PreparedSetData; use prost::Message; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; @@ -51,6 +51,7 @@ use crate::scalar::Scalar; use crate::scalar_fn::Arity; use crate::scalar_fn::ChildName; use crate::scalar_fn::ExecutionArgs; +use crate::scalar_fn::ReduceNode; use crate::scalar_fn::ScalarFnId; use crate::scalar_fn::ScalarFnVTable; use crate::scalar_fn::ScalarFnVTableExt; @@ -207,26 +208,60 @@ impl ScalarFnVTable for ListContains { let value_array = args.get(1)?; // Borrow the list: a constant list scalar owns every element, so cloning it is not free. - if let Some(list) = list_array.as_opt::() + if let Some(list) = constant_list(&list_array) && let Some(value_scalar) = value_array.as_constant() { - let result = compute_contains_scalar(list.scalar(), &value_scalar, options)?; + let result = compute_contains_scalar(list, &value_scalar, options)?; return Ok(ConstantArray::new(result, args.row_count()).into_array()); } compute_list_contains(&list_array, &value_array, options, ctx) } - fn simplify_untyped( + /// A constant list and a constant needle fold to their constant answer, in expression and + /// array trees alike. A [`PreparedSetArray`] list folds through its own parent rule. + fn reduce(&self, options: &Self::Options, node: &T) -> VortexResult> { + let Some(list) = node.child(0).as_constant() else { + return Ok(None); + }; + let Some(needle) = node.child(1).as_constant() else { + return Ok(None); + }; + let result = compute_contains_scalar(&list, &needle, options)?; + Ok(Some(node.new_constant(result))) + } + + /// The validity of `needle IN list` for a literal list, decided without a probe. + /// + /// Off SQL null semantics only a null needle gives null, unless the list is empty, which + /// answers `false` for every needle. Under them a null element can give null too, by leaving + /// a non-match unknown, so only a list without one is decided from the needle alone. A null + /// list gives null for every needle. + /// + /// Without this the validity of a node over a constant list would execute the node, which + /// prepares the list as a set only to read the validity off the result. + fn validity( &self, options: &Self::Options, - expr: &Expression, + expression: &Expression, ) -> VortexResult> { - let Some(list) = expr.child(0).as_opt::() else { + let Some(list) = expression.child(0).as_opt::() else { return Ok(None); }; - Ok(normalized_set(list, options) - .map(|list| ListContains.new_expr(*options, [lit(list), expr.child(1).clone()]))) + if !matches!(list.dtype(), DType::List(..)) { + return Ok(None); + } + let Some(elements) = list.as_list().element_values() else { + return Ok(Some(lit(false))); + }; + + if options.sql_null_semantics && elements.iter().any(Option::is_none) { + return Ok(None); + } + if elements.is_empty() && !options.sql_null_semantics { + return Ok(Some(lit(true))); + } + Ok(Some(expression.child(1).validity()?)) } // Off SQL null semantics an empty list answers `false` even for a null needle; on them a null @@ -240,56 +275,6 @@ impl ScalarFnVTable for ListContains { } } -/// A constant list rewritten into the form a set probe wants — its elements sorted, without -/// duplicates, and with at most one null — or `None` when it is in that form already or its -/// elements have no total order. -/// -/// Neither the order of the elements nor their repetition changes a membership test, and one null -/// element decides as much as many. Whether a null survives does matter: under SQL null semantics -/// it makes a non-match unknown, and off them a list of nothing but nulls is still not empty, which -/// a null needle tells apart. -/// -/// Normalizing once, while the expression is optimized, spares every batch the sort. -fn normalized_set(list: &Scalar, options: &ListContainsOptions) -> Option { - let DType::List(element_dtype, nullability) = list.dtype() else { - return None; - }; - if !matches!( - element_dtype.as_ref(), - DType::Bool(_) - | DType::Primitive(..) - | DType::Decimal(..) - | DType::Utf8(_) - | DType::Binary(_) - ) { - return None; - } - let elements = list.as_list().elements()?; - - let had_null = elements.iter().any(Scalar::is_null); - let mut set: Vec = elements - .iter() - .filter(|element| !element.is_null()) - .cloned() - .collect(); - let mut incomparable = false; - set.sort_by(|a, b| { - a.partial_cmp(b).unwrap_or_else(|| { - incomparable = true; - Ordering::Equal - }) - }); - if incomparable { - return None; - } - set.dedup(); - if had_null && (options.sql_null_semantics || set.is_empty()) { - set.push(Scalar::null(element_dtype.as_ref().clone())); - } - - (set != elements).then(|| Scalar::list(Arc::clone(element_dtype), set, *nullability)) -} - fn compute_contains_scalar( list: &Scalar, needle: &Scalar, @@ -360,13 +345,35 @@ fn compute_list_contains( return list_contains_scalar(array, &value_scalar, nullability, options, ctx); } - if !array.is::() { - return lists_contain_needles(array, value, nullability, options, ctx); + if let Some(set) = array.as_opt::() { + return set.contains(value, options, ctx); } - let set = ListContainsSet::try_new(array, value.dtype(), options, ctx)? - .ok_or_else(|| vortex_err!("A non-null constant list of {} has a set", value.dtype()))?; - set.prepare(ctx)?.contains(value, ctx) + if let Some(list) = array.as_opt::() { + let set = PreparedSetArray::try_new(list.scalar().clone(), array.len(), ctx)?; + + // A canonical needle has no encoding for a kernel to use, so probe it now. + if value.is_canonical() { + return set.data().contains(value, options, ctx); + } + + // Give the prepared set back to the executor. Thus a kernel of the needle encoding can + // probe it, for example on the values of a dictionary only. When no kernel does, the next + // execution finds the prepared set above and probes it, so this happens at most once. + return Ok( + ListContains::try_new_opts(set.into_array(), value.clone(), *options)?.into_array(), + ); + } + + lists_contain_needles(array, value, nullability, options, ctx) +} + +/// The list that every row of `array` holds, when the array is a constant or a prepared set. +fn constant_list(array: &ArrayRef) -> Option<&Scalar> { + if let Some(constant) = array.as_opt::() { + return Some(constant.data().scalar()); + } + array.as_opt::().map(|set| set.data().list()) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -789,6 +796,7 @@ mod tests { use crate::expr::get_item; use crate::expr::gt; use crate::expr::in_list; + use crate::expr::is_not_null; use crate::expr::list_contains; use crate::expr::list_contains_opts; use crate::expr::lit; @@ -796,6 +804,7 @@ mod tests { use crate::expr::or; use crate::expr::root; use crate::expr::stats::Stat; + use crate::optimizer::ArrayOptimizer; use crate::scalar::PValue; use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::ListContains; @@ -1653,81 +1662,6 @@ mod tests { assert_rows_agree(result, BoolArray::from_iter([true, false])) } - fn optimized_set(expr: &Expression) -> VortexResult { - let optimized = expr.optimize_recursive(&DType::Primitive(I32, Nullability::Nullable))?; - Ok(optimized - .child(0) - .as_opt::() - .vortex_expect("the set stays a literal") - .clone()) - } - - #[rstest] - // Off SQL null semantics a null element next to others decides nothing. - #[case::default(ListContainsOptions::default(), vec![Some(3), Some(1), None, Some(3), Some(2)], vec![Some(1), Some(2), Some(3)])] - // On them one null is kept, last, to leave a non-match unknown. - #[case::sql(SQL, vec![Some(3), None, Some(1), None, Some(3)], vec![Some(1), Some(3), None])] - // A list of nothing but nulls keeps one, since an empty list answers a null needle apart. - #[case::only_nulls(ListContainsOptions::default(), vec![None, None], vec![None])] - fn optimize_normalizes_the_set( - #[case] options: ListContainsOptions, - #[case] set: Vec>, - #[case] expected: Vec>, - ) -> VortexResult<()> { - let expr = list_contains_opts(lit(i32_set(set)), root(), options); - assert_eq!(optimized_set(&expr)?, i32_set(expected.clone())); - // A normalized set is left alone, so optimization reaches a fixed point. - let normalized = list_contains_opts(lit(i32_set(expected.clone())), root(), options); - assert_eq!(optimized_set(&normalized)?, i32_set(expected)); - Ok(()) - } - - #[rstest] - #[case::default(ListContainsOptions::default())] - #[case::sql(SQL)] - fn optimize_keeps_the_answer(#[case] options: ListContainsOptions) -> VortexResult<()> { - // Floats compare bitwise: `-0.0` and `0.0` stay distinct members, and NaN is one member. - let element = DType::Primitive(PType::F64, Nullability::Nullable); - let set = Scalar::list( - Arc::new(element.clone()), - [ - Some(2.0), - Some(f64::NAN), - None, - Some(-0.0), - Some(2.0), - Some(f64::NAN), - ] - .into_iter() - .map(|v| match v { - Some(v) => Scalar::primitive(v, Nullability::Nullable), - None => Scalar::null(element.clone()), - }) - .collect(), - Nullability::NonNullable, - ); - let needles = PrimitiveArray::from_option_iter([ - Some(2.0f64), - Some(0.0), - Some(-0.0), - Some(f64::NAN), - Some(5.0), - None, - ]) - .into_array(); - let expr = list_contains_opts(lit(set), root(), options); - let optimized = expr.optimize_recursive(needles.dtype())?; - assert_ne!(optimized, expr, "the set is normalized"); - - let mut ctx = array_session().create_execution_ctx(); - let expected = needles - .clone() - .apply(&expr)? - .execute::(&mut ctx)?; - assert_arrays_eq!(needles.apply(&optimized)?, expected, &mut ctx); - Ok(()) - } - /// Probes `set` with each element, its neighbours and the type's extremes, against a naive /// oracle. fn assert_integer_membership(set: Vec) -> VortexResult<()> @@ -2048,4 +1982,63 @@ mod tests { ); Ok(()) } + + #[rstest] + #[case::default(ListContainsOptions::default(), Some(false))] + #[case::sql(SQL, None)] + fn constant_needle_in_constant_set_folds_to_a_constant( + #[case] options: ListContainsOptions, + #[case] expected: Option, + ) -> VortexResult<()> { + // `3 IN (2, NULL)` is decided at optimization, in an expression and in an array tree. + let set = i32_set(vec![Some(2), None]); + let needle = Scalar::primitive(3i32, Nullability::Nullable); + let expected = match expected { + Some(value) => Scalar::bool(value, Nullability::Nullable), + None => Scalar::null(DType::Bool(Nullability::Nullable)), + }; + + let expr = list_contains_opts(lit(set.clone()), lit(needle.clone()), options) + .bind(&DType::Primitive(I32, Nullability::Nullable))? + .optimize()?; + assert_eq!(expr.as_opt::(), Some(&expected)); + + let array = ListContains::try_new_opts( + ConstantArray::new(set, 3).into_array(), + ConstantArray::new(needle, 3).into_array(), + options, + )? + .into_array() + .optimize()?; + assert_eq!(array.as_constant(), Some(expected)); + Ok(()) + } + + #[test] + fn validity_of_a_constant_set_needs_no_probe() -> VortexResult<()> { + // Off SQL null semantics the needle alone decides validity. Under them a null element + // makes validity depend on the probe, which only the fallback computes. + let needle = col("a"); + let with_null = lit(i32_set(vec![Some(2), None])); + let without_null = lit(i32_set(vec![Some(2)])); + + let expr = list_contains(with_null.clone(), needle.clone()); + assert_eq!(expr.validity()?, needle.validity()?); + let expr = in_list(needle.clone(), without_null); + assert_eq!(expr.validity()?, needle.validity()?); + let expr = in_list(needle.clone(), with_null); + assert_eq!(expr.validity()?, is_not_null(expr)); + + let expr = list_contains(lit(empty_i32_list()), needle.clone()); + assert_eq!(expr.validity()?, lit(true)); + let null_list = Scalar::null(DType::List( + Arc::new(DType::Primitive(I32, Nullability::Nullable)), + Nullability::Nullable, + )); + assert_eq!( + list_contains(lit(null_list), needle).validity()?, + lit(false) + ); + Ok(()) + } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs new file mode 100644 index 00000000000..a1f5b07c30e --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -0,0 +1,387 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use std::fmt::Debug; +use std::fmt::Display; +use std::fmt::Formatter; +use std::hash::Hash; +use std::hash::Hasher; +use std::ops::Range; +use std::sync::Arc; + +use vortex_buffer::BitBuffer; +use vortex_error::VortexExpect; +use vortex_error::VortexResult; +use vortex_error::vortex_bail; +use vortex_error::vortex_ensure; +use vortex_error::vortex_panic; +use vortex_mask::Mask; +use vortex_session::VortexSession; +use vortex_session::registry::CachedId; + +use super::Probe; +use crate::ArrayEq; +use crate::ArrayHash; +use crate::ArrayParts; +use crate::ArrayRef; +use crate::EqMode; +use crate::ExecutionCtx; +use crate::ExecutionResult; +use crate::IntoArray; +use crate::array::Array; +use crate::array::ArrayId; +use crate::array::ArrayView; +use crate::array::OperationsVTable; +use crate::array::VTable; +use crate::array::ValidityVTable; +use crate::array::with_empty_buffers; +use crate::arrays::BoolArray; +use crate::arrays::ConstantArray; +use crate::arrays::ScalarFn; +use crate::arrays::constant::list_scalar_elements; +use crate::arrays::filter::FilterReduce; +use crate::arrays::filter::FilterReduceAdaptor; +use crate::arrays::scalar_fn::ExactScalarFn; +use crate::arrays::scalar_fn::ScalarFnArrayExt; +use crate::arrays::scalar_fn::ScalarFnArrayView; +use crate::arrays::slice::SliceReduce; +use crate::arrays::slice::SliceReduceAdaptor; +use crate::buffer::BufferHandle; +use crate::dtype::DType; +use crate::optimizer::rules::ArrayParentReduceRule; +use crate::optimizer::rules::ParentRuleSet; +use crate::scalar::Scalar; +use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::scalar_fn::fns::list_contains::compute_contains_scalar; +use crate::serde::ArrayChildren; +use crate::validity::Validity; + +/// A [`PreparedSet`]-encoded array. +pub type PreparedSetArray = Array; + +/// A constant, non-null list whose elements are prepared as a set for membership probes. +/// +/// Every row holds the same list, as in a [`ConstantArray`]. When [`ListContains`] gets a constant +/// list and a needle that is not canonical, it puts this array in place of the list and gives the +/// node back to the executor. Thus a [`ListContainsElementKernel`] of the needle encoding gets the +/// prepared set as its list, and can probe its own values with [`PreparedSetData::contains`]. +/// +/// The probe is shared, so a slice or a filter of this array does not build it again, and a +/// constant needle folds against it at optimization. This encoding exists only during execution, +/// and it cannot be serialized. +/// +/// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains +/// [`ListContainsElementKernel`]: crate::scalar_fn::fns::list_contains::ListContainsElementKernel +#[derive(Clone, Debug)] +pub struct PreparedSet; + +/// The data of a [`PreparedSetArray`]: the list, and the probe built from its elements. +#[derive(Clone)] +pub struct PreparedSetData { + list: Scalar, + pub(super) set: Arc, +} + +/// The non-null elements of the list in a probe structure, and the facts about the list that +/// decide the answer when no element matches. +pub(super) struct ElementSet { + pub(super) probe: Probe, + /// Whether the list holds a null element, which a probe does not hold. + has_null_element: bool, + /// Whether the list holds no element at all, counting null elements. + is_empty: bool, +} + +impl PreparedSetData { + /// Prepares the elements of the non-null list scalar `list`. + pub(super) fn try_new(list: Scalar, ctx: &mut ExecutionCtx) -> VortexResult { + vortex_ensure!( + matches!(list.dtype(), DType::List(..)), + "A prepared set needs a list, got {}", + list.dtype() + ); + vortex_ensure!(!list.is_null(), "A prepared set needs a non-null list"); + + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + let is_empty = elements.is_empty(); + + // A null element never equals a needle, so the probe is better off without it. + let valid = elements.validity()?.execute_mask(elements.len(), ctx)?; + let has_null_element = !valid.all_true(); + let elements = if has_null_element { + elements.filter(valid)? + } else { + elements + }; + + let probe = Probe::try_new(elements, ctx)?; + + Ok(Self { + list, + set: Arc::new(ElementSet { + probe, + has_null_element, + is_empty, + }), + }) + } + + /// The list that every row holds. + pub fn list(&self) -> &Scalar { + &self.list + } + + /// Whether each of `needles` is an element of the list, under `options`. + /// + /// The needles must have the dtype of the list's elements, ignoring nullability. The result + /// has one row per needle, and the nullability that [`ListContainsOptions::result_nullability`] + /// declares. A null needle gives `null`. The exception is an empty list off SQL null + /// semantics, which gives `false` for every needle. Under SQL null semantics, a list that holds + /// a null element gives `null` for a needle that matches no element. + pub fn contains( + &self, + needles: &ArrayRef, + options: &ListContainsOptions, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let DType::List(element_dtype, _) = self.list.dtype() else { + vortex_panic!("A prepared set always holds a list"); + }; + if !element_dtype.eq_ignore_nullability(needles.dtype()) { + vortex_bail!( + "Element type {} of list does not match search value {}", + element_dtype, + needles.dtype(), + ); + } + + let (bits, needle_validity) = self.set.probe.contains(needles, ctx)?; + self.finish(bits, needle_validity, needles.dtype(), options) + } + + /// Makes the result from one membership bit per needle and the validity of the needles. + fn finish( + &self, + bits: BitBuffer, + needle_validity: Validity, + needle_dtype: &DType, + options: &ListContainsOptions, + ) -> VortexResult { + let nullability = options.result_nullability(self.list.dtype(), needle_dtype); + + let validity = if self.set.is_empty && !options.sql_null_semantics { + Validity::NonNullable + } else if options.sql_null_semantics && self.set.has_null_element { + // Only a match is known. A comparison with the null element makes a non-match unknown. + needle_validity.and(Validity::from(bits.clone()))? + } else { + needle_validity + }; + + Ok(BoolArray::new(bits, validity.union_nullability(nullability)).into_array()) + } +} + +impl Debug for PreparedSetData { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PreparedSetData") + .field("list", &self.list) + .finish_non_exhaustive() + } +} + +impl Display for PreparedSetData { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "list: {}", self.list) + } +} + +impl ArrayHash for PreparedSetData { + fn array_hash(&self, state: &mut H, _accuracy: EqMode) { + self.list.hash(state); + } +} + +impl ArrayEq for PreparedSetData { + fn array_eq(&self, other: &Self, _accuracy: EqMode) -> bool { + self.list == other.list + } +} + +impl Array { + /// Prepares the non-null list scalar `list` as a set, repeated `len` times. + pub(crate) fn try_new(list: Scalar, len: usize, ctx: &mut ExecutionCtx) -> VortexResult { + let data = PreparedSetData::try_new(list, ctx)?; + Ok(Self::from_data(data, len)) + } + + /// An array of `len` rows that share the prepared set `data`. + fn from_data(data: PreparedSetData, len: usize) -> Self { + let dtype = data.list.dtype().clone(); + + // SAFETY: the dtype is the dtype of the list that every row holds. + unsafe { Array::from_parts_unchecked(ArrayParts::new(PreparedSet, dtype, len, data)) } + } +} + +const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ + ParentRuleSet::lift(&ConstantNeedleRule), + ParentRuleSet::lift(&FilterReduceAdaptor(PreparedSet)), + ParentRuleSet::lift(&SliceReduceAdaptor(PreparedSet)), +]); + +/// Folds `list_contains` of a constant needle against the prepared set into its constant answer. +/// +/// [`ListContains::reduce`] folds a [`ConstantArray`] list the same way. This rule covers the list +/// after it is prepared, which the generic constant check does not see. +#[derive(Debug)] +struct ConstantNeedleRule; + +impl ArrayParentReduceRule for ConstantNeedleRule { + type Parent = ExactScalarFn; + + fn reduce_parent( + &self, + array: ArrayView<'_, PreparedSet>, + parent: ScalarFnArrayView<'_, ListContains>, + child_idx: usize, + ) -> VortexResult> { + // The prepared set is the list child. As the needle it is a list of lists, which this rule + // does not fold. + if child_idx != 0 { + return Ok(None); + } + let scalar_fn_array = parent + .as_opt::() + .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); + let Some(needle) = scalar_fn_array.get_child(1).as_constant() else { + return Ok(None); + }; + + let result = compute_contains_scalar(array.list(), &needle, parent.options)?; + Ok(Some( + ConstantArray::new(result, scalar_fn_array.len()).into_array(), + )) + } +} + +impl VTable for PreparedSet { + type TypedArrayData = PreparedSetData; + + type OperationsVTable = Self; + type ValidityVTable = Self; + + fn id(&self) -> ArrayId { + static ID: CachedId = CachedId::new("vortex.list.prepared_set"); + *ID + } + + fn validate( + &self, + data: &PreparedSetData, + dtype: &DType, + _len: usize, + _slots: &[Option], + ) -> VortexResult<()> { + vortex_ensure!( + data.list.dtype() == dtype, + "PreparedSetArray list dtype does not match outer dtype" + ); + Ok(()) + } + + fn nbuffers(_array: ArrayView<'_, Self>) -> usize { + 0 + } + + fn buffer(_array: ArrayView<'_, Self>, idx: usize) -> BufferHandle { + vortex_panic!("PreparedSetArray buffer index {idx} out of bounds") + } + + fn buffer_name(_array: ArrayView<'_, Self>, _idx: usize) -> Option { + None + } + + fn with_buffers( + &self, + array: ArrayView<'_, Self>, + buffers: &[BufferHandle], + ) -> VortexResult> { + with_empty_buffers(self, array, buffers) + } + + fn slot_name(_array: ArrayView<'_, Self>, idx: usize) -> String { + vortex_panic!("PreparedSetArray slot_name index {idx} out of bounds") + } + + fn serialize( + _array: ArrayView<'_, Self>, + _session: &VortexSession, + ) -> VortexResult>> { + vortex_bail!("PreparedSetArray is not serializable") + } + + fn deserialize( + &self, + _dtype: &DType, + _len: usize, + _metadata: &[u8], + _buffers: &[BufferHandle], + _children: &dyn ArrayChildren, + _session: &VortexSession, + ) -> VortexResult> { + vortex_bail!("PreparedSetArray is not serializable") + } + + fn execute(array: Array, _ctx: &mut ExecutionCtx) -> VortexResult { + // The rows hold the list only. The probe is of no use to the canonical form. + Ok(ExecutionResult::done(ConstantArray::new( + array.data().list.clone(), + array.len(), + ))) + } + + fn reduce_parent( + array: ArrayView<'_, Self>, + parent: &ArrayRef, + child_idx: usize, + ) -> VortexResult> { + PARENT_RULES.evaluate(array, parent, child_idx) + } +} + +impl OperationsVTable for PreparedSet { + type ProbeState = (); + + fn scalar_at( + array: ArrayView<'_, PreparedSet>, + _index: usize, + _ctx: &mut ExecutionCtx, + ) -> VortexResult { + Ok(array.list.clone()) + } +} + +impl ValidityVTable for PreparedSet { + fn validity(_array: ArrayView<'_, PreparedSet>) -> VortexResult { + // The list is never null. + Ok(Validity::AllValid) + } +} + +impl SliceReduce for PreparedSet { + fn slice(array: ArrayView<'_, Self>, range: Range) -> VortexResult> { + Ok(Some( + PreparedSetArray::from_data(array.data().clone(), range.len()).into_array(), + )) + } +} + +impl FilterReduce for PreparedSet { + fn filter(array: ArrayView<'_, Self>, mask: &Mask) -> VortexResult> { + Ok(Some( + PreparedSetArray::from_data(array.data().clone(), mask.true_count()).into_array(), + )) + } +} diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs index e87690d1c3f..900c7ee0e0d 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -1,11 +1,16 @@ // SPDX-License-Identifier: Apache-2.0 // SPDX-FileCopyrightText: Copyright the Vortex contributors -//! Probing a constant set: the probe structure is built from a [`ListContainsSet`] for canonical -//! evaluation. Chunked expressions dispatch per chunk so each encoding can specialize membership. +//! Probing a constant set. [`PreparedSetArray`] holds a constant list together with the probe +//! built from its elements, so that the kernels of a needle encoding can probe their own values. + +mod array; use std::hash::BuildHasher; +pub use array::PreparedSet; +pub use array::PreparedSetArray; +pub use array::PreparedSetData; use num_traits::ToPrimitive; use num_traits::WrappingSub; use vortex_buffer::BitBuffer; @@ -39,7 +44,6 @@ use crate::match_each_integer_ptype; use crate::scalar::DecimalValue; use crate::scalar_fn::fns::binary::build_row_comparator; use crate::scalar_fn::fns::binary::collect_bits; -use crate::scalar_fn::fns::list_contains::ListContainsSet; use crate::validity::Validity; /// A set whose span of values needs at most this many bits per element is probed through a bitmap @@ -48,18 +52,12 @@ const BITMAP_BITS_PER_ELEMENT: u128 = 64; /// A span this narrow is probed through a bitmap whatever the size of the set. const BITMAP_MIN_BITS: u128 = 1 << 12; -/// A [`ListContainsSet`] with the structure that probes it, prepared once and then run over any -/// number of needle arrays of the set's dtype. +/// The structure that probes the non-null elements of a set. /// /// Every needle is probed against a bitmap, a sorted set or a hash table. Decimals use their /// unscaled integers with the same bitmap and sorted-value strategy as primitives. Nested /// values use sorted row indices with the same comparator as equality. No probe constructs /// per-element expressions or materializes scalars in its loop. -pub struct PreparedSet { - set: ListContainsSet, - probe: Probe, -} - enum Probe { /// Integers, or floats by their bit patterns, spanning a dense range: one bit per value of the /// span above `min_offset`, the smallest element as a `usize`. @@ -91,15 +89,13 @@ enum Probe { }, } -impl ListContainsSet { - /// Builds the structure that probes this set. - pub fn prepare(self, ctx: &mut ExecutionCtx) -> VortexResult { - let probe = match self.elements().dtype() { +impl Probe { + /// Builds the structure that probes the non-null `elements`. + fn try_new(elements: ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { + Ok(match elements.dtype() { DType::Primitive(ptype, _) => { let ptype = bit_pattern_ptype(*ptype); - let elements = self - .elements() - .clone() + let elements = elements .execute::(ctx)? .reinterpret_cast(ptype); match_each_integer_ptype!(ptype, |T| { @@ -107,7 +103,7 @@ impl ListContainsSet { }) } DType::Decimal(..) => { - let elements = self.elements().clone().execute::(ctx)?; + let elements = elements.execute::(ctx)?; match_each_decimal_value_type!(elements.values_type(), |T| { let values = elements.buffer::(); if let Some((min, bitmap)) = integer_bitmap(&values, ctx.allocator()) { @@ -125,15 +121,10 @@ impl ListContainsSet { }) } DType::Utf8(_) | DType::Binary(_) => { - bytes_probe(self.elements().clone().execute::(ctx)?) + bytes_probe(elements.execute::(ctx)?) } _ => { - let elements = self - .elements() - .clone() - .execute::(ctx)? - .0 - .into_array(); + let elements = elements.execute::(ctx)?.0.into_array(); let mut indices: Vec = (0..elements.len()).collect(); if !indices.is_empty() { let compare = build_row_comparator(&elements, &elements, ctx)?; @@ -144,22 +135,17 @@ impl ListContainsSet { } Probe::Rows { elements, indices } } - }; - Ok(PreparedSet { set: self, probe }) - } -} - -impl PreparedSet { - /// The dtype of every result, as the expression declares it. - pub fn dtype(&self) -> DType { - DType::Bool(self.set.nullability) + }) } - /// Whether each of `needles`, which have the set's dtype, is a member of the set. - /// - /// The result has the nullability the expression declares, whatever the needles' encoding. - pub fn contains(&self, needles: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { - match &self.probe { + /// One membership bit per needle, and the validity of the needles, which have the dtype of the + /// elements. + fn contains( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult<(BitBuffer, Validity)> { + match self { Probe::Bitmap { .. } | Probe::Sorted(_) => self.contains_primitive(needles, ctx), Probe::DecimalBitmap { .. } | Probe::DecimalSorted(_) => { self.contains_decimal(needles, ctx) @@ -168,9 +154,9 @@ impl PreparedSet { elements, hasher, table, - } => self.contains_bytes(elements, hasher, table, needles, ctx), + } => Self::contains_bytes(elements, hasher, table, needles, ctx), Probe::Rows { elements, indices } => { - self.contains_rows(elements, indices, needles, ctx) + Self::contains_rows(elements, indices, needles, ctx) } } } @@ -182,26 +168,25 @@ impl PreparedSet { &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, - ) -> VortexResult { + ) -> VortexResult<(BitBuffer, Validity)> { let primitive = needles.clone().execute::(ctx)?; let ptype = bit_pattern_ptype(primitive.ptype()); let values = primitive.reinterpret_cast(ptype); let bits = match_each_integer_ptype!(ptype, |T| { - self.probe - .integer_bits(values.as_slice::(), ctx.allocator()) + self.integer_bits(values.as_slice::(), ctx.allocator()) }); - self.set.finish(bits, primitive.validity()?) + Ok((bits, primitive.validity()?)) } fn contains_decimal( &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, - ) -> VortexResult { + ) -> VortexResult<(BitBuffer, Validity)> { let needles = needles.clone().execute::(ctx)?; // Logical precision and scale agree, but physical widths can differ. Widen once to // the common width so out-of-range needles cannot truncate into false matches. - let bits = match &self.probe { + let bits = match self { Probe::DecimalBitmap { min, bitmap } => { let common = min.decimal_type().max(needles.values_type()); match_each_decimal_value_type!(common, |T| { @@ -210,9 +195,9 @@ impl PreparedSet { collect_bits( &values, |value| { - value.offset_from(min).is_some_and(|offset| { - offset < bitmap.len() && bitmap.value(offset) - }) + value + .offset_from(min) + .is_some_and(|offset| offset < bitmap.len() && bitmap.value(offset)) }, ctx.allocator(), ) @@ -228,17 +213,16 @@ impl PreparedSet { } _ => unreachable!("decimal needles meet a decimal probe"), }; - self.set.finish(bits, needles.validity()?) + Ok((bits, needles.validity()?)) } fn contains_bytes( - &self, elements: &VarBinViewArray, hasher: &RandomState, table: &HashTable, needles: &ArrayRef, ctx: &mut ExecutionCtx, - ) -> VortexResult { + ) -> VortexResult<(BitBuffer, Validity)> { let element_views = elements.views(); let element_buffers = data_buffers(elements); let array = needles.clone().execute::(ctx)?; @@ -255,21 +239,20 @@ impl PreparedSet { }, ctx.allocator(), ); - self.set.finish(bits, array.validity()?) + Ok((bits, array.validity()?)) } fn contains_rows( - &self, elements: &ArrayRef, indices: &[usize], needles: &ArrayRef, ctx: &mut ExecutionCtx, - ) -> VortexResult { + ) -> VortexResult<(BitBuffer, Validity)> { if indices.is_empty() { - return self.set.finish( + return Ok(( BitBuffer::full_in(false, needles.len(), ctx.allocator().clone()), needles.validity()?, - ); + )); } let needles = needles .clone() @@ -288,11 +271,9 @@ impl PreparedSet { }, ctx.allocator().clone(), ); - self.set.finish(bits, needles.validity()?) + Ok((bits, needles.validity()?)) } -} -impl Probe { /// One bit per needle, set when the needle is an element. fn integer_bits( &self, diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs index c0faa84440a..c08511979a3 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -6,6 +6,8 @@ use vortex_buffer::BitBuffer; use vortex_buffer::buffer; use vortex_error::VortexResult; +use super::PreparedSet; +use super::PreparedSetArray; use super::Probe; use crate::ArrayRef; use crate::IntoArray; @@ -28,16 +30,17 @@ use crate::dtype::Nullability; use crate::dtype::PType; use crate::dtype::i256; use crate::match_each_decimal_value_type; +use crate::optimizer::ArrayOptimizer; use crate::scalar::DecimalValue; use crate::scalar::Scalar; +use crate::scalar_fn::fns::list_contains::ListContains; use crate::scalar_fn::fns::list_contains::ListContainsOptions; -use crate::scalar_fn::fns::list_contains::ListContainsSet; +use crate::scalar_fn::fns::list_contains::PreparedSetData; use crate::validity::Validity; fn nested_needles() -> ArrayRef { ListArray::try_new( - PrimitiveArray::from_option_iter([Some(1i32), None, Some(2), None, Some(9)]) - .into_array(), + PrimitiveArray::from_option_iter([Some(1i32), None, Some(2), None, Some(9)]).into_array(), buffer![0u32, 2, 2, 4, 5].into_array(), Validity::from(BitBuffer::from_iter([true, false, true, true])), ) @@ -98,16 +101,10 @@ fn test_row_set_returns_membership_bits( .collect::>>()?; let mut elements = members.clone(); elements.push(Scalar::null(dtype.clone())); - let list = ConstantArray::new( - Scalar::list(dtype, elements, Nullability::NonNullable), - needles.len(), - ) - .into_array(); + let list = Scalar::list(dtype, elements, Nullability::NonNullable); let options = ListContainsOptions { sql_null_semantics }; - let set = ListContainsSet::try_new(&list, needles.dtype(), &options, &mut ctx)? - .unwrap() - .prepare(&mut ctx)?; - let result = set.contains(&needles, &mut ctx)?; + let set = PreparedSetData::try_new(list, &mut ctx)?; + let result = set.contains(&needles, &options, &mut ctx)?; // The old fallback returned a lazy OR tree here. assert!(result.is::()); let expected = (0..needles.len()) @@ -124,7 +121,21 @@ fn test_row_set_returns_membership_bits( }) }) .collect::>>()?; - assert_arrays_eq!(result, BoolArray::from_iter(expected), &mut ctx); + + // The list is not null, so only a null needle can make the result null. Under SQL null + // semantics the null element also can, by leaving a non-match unknown. + let nullability = needles.dtype().nullability() | Nullability::from(sql_null_semantics); + assert_eq!(result.dtype(), &DType::Bool(nullability)); + + let validity = match nullability { + Nullability::NonNullable => Validity::NonNullable, + Nullability::Nullable => Validity::from_iter(expected.iter().map(Option::is_some)), + }; + let expected = BoolArray::new( + BitBuffer::from_iter(expected.into_iter().map(Option::unwrap_or_default)), + validity, + ); + assert_arrays_eq!(result, expected, &mut ctx); Ok(()) } @@ -132,8 +143,12 @@ fn test_row_set_returns_membership_bits( fn test_decimal_bitmap_across_storage_widths( #[values(2, 4, 9, 18, 38, 76)] precision: u8, #[values( - DecimalType::I8, DecimalType::I16, DecimalType::I32, - DecimalType::I64, DecimalType::I128, DecimalType::I256, + DecimalType::I8, + DecimalType::I16, + DecimalType::I32, + DecimalType::I64, + DecimalType::I128, + DecimalType::I256 )] needle_width: DecimalType, #[values(false, true)] sql_null_semantics: bool, @@ -145,28 +160,21 @@ fn test_decimal_bitmap_across_storage_widths( .map(|value| Scalar::decimal(value.into(), decimal, Nullability::Nullable)) .into(); elements.push(Scalar::null(dtype.clone())); - let list = ConstantArray::new( - Scalar::list(dtype, elements, Nullability::NonNullable), - 4, - ) - .into_array(); + let list = Scalar::list(dtype, elements, Nullability::NonNullable); let needles = match_each_decimal_value_type!(needle_width, |T| { DecimalArray::from_option_iter::( - [Some(-2i8), Some(0), Some(1), None].map(|value| { - value.map(|value| DecimalValue::from(value).cast::().unwrap()) - }), + [Some(-2i8), Some(0), Some(1), None] + .map(|value| value.map(|value| DecimalValue::from(value).cast::().unwrap())), decimal, ) .into_array() }); let options = ListContainsOptions { sql_null_semantics }; - let set = ListContainsSet::try_new(&list, needles.dtype(), &options, &mut ctx)? - .unwrap() - .prepare(&mut ctx)?; - assert!(matches!(&set.probe, Probe::DecimalBitmap { .. })); + let set = PreparedSetData::try_new(list, &mut ctx)?; + assert!(matches!(&set.set.probe, Probe::DecimalBitmap { .. })); let non_match = (!sql_null_semantics).then_some(false); assert_arrays_eq!( - set.contains(&needles, &mut ctx)?, + set.contains(&needles, &options, &mut ctx)?, BoolArray::from_iter([Some(true), non_match, Some(true), None]), &mut ctx ); @@ -186,17 +194,13 @@ fn test_decimal_wide_values( let mut ctx = array_session().create_execution_ctx(); let decimal = DecimalDType::new(precision, 2); let second = if dense { base + i256::ONE } else { -base }; - let list = ConstantArray::new( - Scalar::list( - DType::Decimal(decimal, Nullability::NonNullable), - [second, base, second] - .map(|v| Scalar::decimal(v.into(), decimal, Nullability::NonNullable)) - .into(), - Nullability::NonNullable, - ), - 6, - ) - .into_array(); + let list = Scalar::list( + DType::Decimal(decimal, Nullability::NonNullable), + [second, base, second] + .map(|v| Scalar::decimal(v.into(), decimal, Nullability::NonNullable)) + .into(), + Nullability::NonNullable, + ); // The last non-null value shares the low 64 bits of a member but must not match it. let needles = DecimalArray::from_option_iter::( [ @@ -210,22 +214,91 @@ fn test_decimal_wide_values( decimal, ) .into_array(); - let set = ListContainsSet::try_new( - &list, - needles.dtype(), - &ListContainsOptions::default(), - &mut ctx, - )? - .unwrap() - .prepare(&mut ctx)?; - assert_eq!(matches!(&set.probe, Probe::DecimalBitmap { .. }), dense); - assert_eq!(matches!(&set.probe, Probe::DecimalSorted(_)), !dense); + let set = PreparedSetData::try_new(list, &mut ctx)?; + assert_eq!(matches!(&set.set.probe, Probe::DecimalBitmap { .. }), dense); + assert_eq!(matches!(&set.set.probe, Probe::DecimalSorted(_)), !dense); assert_arrays_eq!( - set.contains(&needles, &mut ctx)?, + set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, BoolArray::from_iter([ - Some(true), Some(true), Some(false), Some(false), Some(false), None, + Some(true), + Some(true), + Some(false), + Some(false), + Some(false), + None, ]), &mut ctx ); Ok(()) } + +/// The set `{2, null}` of nullable `i32`. +fn set_with_null() -> Scalar { + let element = DType::Primitive(PType::I32, Nullability::Nullable); + Scalar::list( + element.clone(), + vec![ + Scalar::primitive(2i32, Nullability::Nullable), + Scalar::null(element), + ], + Nullability::NonNullable, + ) +} + +#[test] +fn test_prepared_set_rows_are_the_constant_list() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + + assert_arrays_eq!(set, ConstantArray::new(set_with_null(), 4), &mut ctx); + + // A slice keeps the probe instead of building it again. + let sliced = set.slice(1..3)?; + assert!(sliced.is::()); + assert_eq!(sliced.len(), 2); + Ok(()) +} + +#[rstest] +#[case::default(ListContainsOptions::default(), [Some(false), Some(true), None, Some(false)])] +#[case::sql( + ListContainsOptions { sql_null_semantics: true }, + [None, Some(true), None, None] +)] +fn test_list_contains_probes_a_prepared_set_list( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 4], +) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + let needles = + PrimitiveArray::from_option_iter([Some(1i32), Some(2), None, Some(3)]).into_array(); + + let result = ListContains::try_new_opts(set, needles, options)?.into_array(); + assert_arrays_eq!(result, BoolArray::from_iter(expected), &mut ctx); + Ok(()) +} + +#[rstest] +#[case::default(ListContainsOptions::default(), Some(false))] +#[case::sql(ListContainsOptions { sql_null_semantics: true }, None)] +fn test_constant_needle_folds_without_probing( + #[case] options: ListContainsOptions, + #[case] expected: Option, +) -> VortexResult<()> { + // `3 IN (2, NULL)` against the prepared set is decided at optimization, as it is against a + // constant list. + let mut ctx = array_session().create_execution_ctx(); + let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + let needle = ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 4).into_array(); + + let optimized = ListContains::try_new_opts(set, needle, options)? + .into_array() + .optimize()?; + let expected = match expected { + Some(value) => Scalar::bool(value, Nullability::Nullable), + None => Scalar::null(DType::Bool(Nullability::Nullable)), + }; + assert_eq!(optimized.as_constant(), Some(expected)); + Ok(()) +} From c0f7b8b89551bb9412eb7ea891b3e6a093c34065 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 18:28:12 -0400 Subject: [PATCH 4/9] fixes Signed-off-by: Robert Kruszewski --- vortex-array/benches/list_contains_set.rs | 2 +- .../src/scalar_fn/fns/list_contains/prepared/array.rs | 5 +++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/vortex-array/benches/list_contains_set.rs b/vortex-array/benches/list_contains_set.rs index 87b2057f9c2..76e75f33dcf 100644 --- a/vortex-array/benches/list_contains_set.rs +++ b/vortex-array/benches/list_contains_set.rs @@ -60,7 +60,7 @@ fn random_i64(len: usize) -> (Vec, Vec) { fn bench_in_set(bencher: Bencher, set: Scalar, needles: ArrayRef) { let session = vortex_array::array_session(); - // Optimized as a scan optimizes it, so the set arrives normalized. + // Optimized as a scan optimizes it. let expr = list_contains(lit(set), root()) .bind(needles.dtype()) .unwrap() diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs index a1f5b07c30e..bf2535364f6 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -233,8 +233,9 @@ const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ /// Folds `list_contains` of a constant needle against the prepared set into its constant answer. /// -/// [`ListContains::reduce`] folds a [`ConstantArray`] list the same way. This rule covers the list -/// after it is prepared, which the generic constant check does not see. +/// [`ListContains`] folds a [`ConstantArray`] list the same way in its +/// [`reduce`](crate::scalar_fn::ScalarFnVTable::reduce). This rule covers the list after it is +/// prepared, which the generic constant check does not see. #[derive(Debug)] struct ConstantNeedleRule; From b91215a40bb62b20de15ccb6fb0f32e8fc5516b1 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 22:11:20 -0400 Subject: [PATCH 5/9] less Signed-off-by: Robert Kruszewski --- .../src/scalar_fn/fns/list_contains/mod.rs | 27 +- .../fns/list_contains/prepared/array.rs | 200 ++++---- .../fns/list_contains/prepared/mod.rs | 461 ++++++++---------- .../fns/list_contains/prepared/tests.rs | 103 +++- 4 files changed, 411 insertions(+), 380 deletions(-) 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 d23906a0b06..0759674a902 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -37,6 +37,7 @@ use crate::arrays::ListViewArray; use crate::arrays::PrimitiveArray; use crate::arrays::ScalarFnArray; use crate::arrays::bool::BoolArrayExt; +use crate::arrays::constant::list_scalar_elements; use crate::arrays::listview::ListViewArraySlotsExt; use crate::arrays::primitive::PrimitiveArrayExt; use crate::dtype::DType; @@ -208,10 +209,10 @@ impl ScalarFnVTable for ListContains { let value_array = args.get(1)?; // Borrow the list: a constant list scalar owns every element, so cloning it is not free. - if let Some(list) = constant_list(&list_array) + if let Some(list) = list_array.as_opt::() && let Some(value_scalar) = value_array.as_constant() { - let result = compute_contains_scalar(list, &value_scalar, options)?; + let result = compute_contains_scalar(list.scalar(), &value_scalar, options)?; return Ok(ConstantArray::new(result, args.row_count()).into_array()); } @@ -219,7 +220,7 @@ impl ScalarFnVTable for ListContains { } /// A constant list and a constant needle fold to their constant answer, in expression and - /// array trees alike. A [`PreparedSetArray`] list folds through its own parent rule. + /// array trees alike. A [`PreparedSetArray`] list looks a constant needle up at execution. fn reduce(&self, options: &Self::Options, node: &T) -> VortexResult> { let Some(list) = node.child(0).as_constant() else { return Ok(None); @@ -339,18 +340,20 @@ fn compute_list_contains( .into_array()); } + if let Some(set) = array.as_opt::() { + return set.contains(value, options, ctx); + } + let nullability = options.result_nullability(array.dtype(), value.dtype()); if let Some(value_scalar) = value.as_constant() { return list_contains_scalar(array, &value_scalar, nullability, options, ctx); } - if let Some(set) = array.as_opt::() { - return set.contains(value, options, ctx); - } - if let Some(list) = array.as_opt::() { - let set = PreparedSetArray::try_new(list.scalar().clone(), array.len(), ctx)?; + let elements = list_scalar_elements(&list.scalar().as_list(), ctx.allocator()); + let set = + PreparedSetArray::try_new(elements, array.dtype().nullability(), array.len(), ctx)?; // A canonical needle has no encoding for a kernel to use, so probe it now. if value.is_canonical() { @@ -368,14 +371,6 @@ fn compute_list_contains( lists_contain_needles(array, value, nullability, options, ctx) } -/// The list that every row of `array` holds, when the array is a constant or a prepared set. -fn constant_list(array: &ArrayRef) -> Option<&Scalar> { - if let Some(constant) = array.as_opt::() { - return Some(constant.data().scalar()); - } - array.as_opt::().map(|set| set.data().list()) -} - /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. /// /// This is the canonical implementation, over an executed [`ListViewArray`]. diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs index bf2535364f6..d276fbe10de 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -20,6 +20,7 @@ use vortex_session::VortexSession; use vortex_session::registry::CachedId; use super::Probe; +use super::new_probe; use crate::ArrayEq; use crate::ArrayHash; use crate::ArrayParts; @@ -37,23 +38,18 @@ use crate::array::ValidityVTable; use crate::array::with_empty_buffers; use crate::arrays::BoolArray; use crate::arrays::ConstantArray; -use crate::arrays::ScalarFn; -use crate::arrays::constant::list_scalar_elements; +use crate::arrays::ListViewArray; use crate::arrays::filter::FilterReduce; use crate::arrays::filter::FilterReduceAdaptor; -use crate::arrays::scalar_fn::ExactScalarFn; -use crate::arrays::scalar_fn::ScalarFnArrayExt; -use crate::arrays::scalar_fn::ScalarFnArrayView; use crate::arrays::slice::SliceReduce; use crate::arrays::slice::SliceReduceAdaptor; use crate::buffer::BufferHandle; use crate::dtype::DType; -use crate::optimizer::rules::ArrayParentReduceRule; +use crate::dtype::Nullability; +use crate::match_smallest_list_offset_type; use crate::optimizer::rules::ParentRuleSet; use crate::scalar::Scalar; -use crate::scalar_fn::fns::list_contains::ListContains; use crate::scalar_fn::fns::list_contains::ListContainsOptions; -use crate::scalar_fn::fns::list_contains::compute_contains_scalar; use crate::serde::ArrayChildren; use crate::validity::Validity; @@ -67,26 +63,28 @@ pub type PreparedSetArray = Array; /// node back to the executor. Thus a [`ListContainsElementKernel`] of the needle encoding gets the /// prepared set as its list, and can probe its own values with [`PreparedSetData::contains`]. /// -/// The probe is shared, so a slice or a filter of this array does not build it again, and a -/// constant needle folds against it at optimization. This encoding exists only during execution, -/// and it cannot be serialized. +/// The probe is shared, so a slice or a filter of this array does not build it again. This encoding +/// exists only during execution, and it cannot be serialized. /// /// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains /// [`ListContainsElementKernel`]: crate::scalar_fn::fns::list_contains::ListContainsElementKernel #[derive(Clone, Debug)] pub struct PreparedSet; -/// The data of a [`PreparedSetArray`]: the list, and the probe built from its elements. +/// The data of a [`PreparedSetArray`]: the elements of the list, and the probe built from them. #[derive(Clone)] pub struct PreparedSetData { - list: Scalar, + /// All elements of the list, null elements included, with the element dtype of the list. + elements: ArrayRef, + /// The nullability of the list dtype. The list itself is never null. + nullability: Nullability, pub(super) set: Arc, } /// The non-null elements of the list in a probe structure, and the facts about the list that /// decide the answer when no element matches. pub(super) struct ElementSet { - pub(super) probe: Probe, + pub(super) probe: Box, /// Whether the list holds a null element, which a probe does not hold. has_null_element: bool, /// Whether the list holds no element at all, counting null elements. @@ -94,31 +92,29 @@ pub(super) struct ElementSet { } impl PreparedSetData { - /// Prepares the elements of the non-null list scalar `list`. - pub(super) fn try_new(list: Scalar, ctx: &mut ExecutionCtx) -> VortexResult { - vortex_ensure!( - matches!(list.dtype(), DType::List(..)), - "A prepared set needs a list, got {}", - list.dtype() - ); - vortex_ensure!(!list.is_null(), "A prepared set needs a non-null list"); - - let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + /// Prepares `elements`, the elements of a non-null list with the list nullability + /// `nullability`. + pub(super) fn try_new( + elements: ArrayRef, + nullability: Nullability, + ctx: &mut ExecutionCtx, + ) -> VortexResult { let is_empty = elements.is_empty(); // A null element never equals a needle, so the probe is better off without it. let valid = elements.validity()?.execute_mask(elements.len(), ctx)?; let has_null_element = !valid.all_true(); - let elements = if has_null_element { + let probe_elements = if has_null_element { elements.filter(valid)? } else { - elements + elements.clone() }; - let probe = Probe::try_new(elements, ctx)?; + let probe = new_probe(probe_elements, ctx)?; Ok(Self { - list, + elements, + nullability, set: Arc::new(ElementSet { probe, has_null_element, @@ -127,9 +123,14 @@ impl PreparedSetData { }) } - /// The list that every row holds. - pub fn list(&self) -> &Scalar { - &self.list + /// The elements of the list that every row holds. + pub fn elements(&self) -> &ArrayRef { + &self.elements + } + + /// The dtype of the list that every row holds. + fn list_dtype(&self) -> DType { + DType::List(Arc::new(self.elements.dtype().clone()), self.nullability) } /// Whether each of `needles` is an element of the list, under `options`. @@ -139,15 +140,15 @@ impl PreparedSetData { /// declares. A null needle gives `null`. The exception is an empty list off SQL null /// semantics, which gives `false` for every needle. Under SQL null semantics, a list that holds /// a null element gives `null` for a needle that matches no element. + /// + /// A constant needle is probed once, and gives a constant result. pub fn contains( &self, needles: &ArrayRef, options: &ListContainsOptions, ctx: &mut ExecutionCtx, ) -> VortexResult { - let DType::List(element_dtype, _) = self.list.dtype() else { - vortex_panic!("A prepared set always holds a list"); - }; + let element_dtype = self.elements.dtype(); if !element_dtype.eq_ignore_nullability(needles.dtype()) { vortex_bail!( "Element type {} of list does not match search value {}", @@ -156,10 +157,31 @@ impl PreparedSetData { ); } + if let Some(needle) = needles.as_constant() { + let result = self.contains_scalar(&needle, options, ctx)?; + return Ok(ConstantArray::new(result, needles.len()).into_array()); + } + let (bits, needle_validity) = self.set.probe.contains(needles, ctx)?; self.finish(bits, needle_validity, needles.dtype(), options) } + /// Whether the constant `needle` is an element of the list, under `options`. + /// + /// The answer is the one [`Self::contains`] gives for each row of a needle with this value. The + /// probe gets one row, not a row per needle. + fn contains_scalar( + &self, + needle: &Scalar, + options: &ListContainsOptions, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let needles = ConstantArray::new(needle.clone(), 1).into_array(); + let (bits, needle_validity) = self.set.probe.contains(&needles, ctx)?; + self.finish(bits, needle_validity, needle.dtype(), options)? + .execute_scalar(0, ctx) + } + /// Makes the result from one membership bit per needle and the validity of the needles. fn finish( &self, @@ -168,7 +190,7 @@ impl PreparedSetData { needle_dtype: &DType, options: &ListContainsOptions, ) -> VortexResult { - let nullability = options.result_nullability(self.list.dtype(), needle_dtype); + let nullability = options.result_nullability(&self.list_dtype(), needle_dtype); let validity = if self.set.is_empty && !options.sql_null_semantics { Validity::NonNullable @@ -186,39 +208,47 @@ impl PreparedSetData { impl Debug for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { f.debug_struct("PreparedSetData") - .field("list", &self.list) + .field("elements", &self.elements) + .field("nullability", &self.nullability) .finish_non_exhaustive() } } impl Display for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - write!(f, "list: {}", self.list) + write!(f, "elements: {}", self.elements.len()) } } impl ArrayHash for PreparedSetData { - fn array_hash(&self, state: &mut H, _accuracy: EqMode) { - self.list.hash(state); + fn array_hash(&self, state: &mut H, accuracy: EqMode) { + self.elements.array_hash(state, accuracy); + self.nullability.hash(state); } } impl ArrayEq for PreparedSetData { - fn array_eq(&self, other: &Self, _accuracy: EqMode) -> bool { - self.list == other.list + fn array_eq(&self, other: &Self, accuracy: EqMode) -> bool { + self.nullability == other.nullability && self.elements.array_eq(&other.elements, accuracy) } } impl Array { - /// Prepares the non-null list scalar `list` as a set, repeated `len` times. - pub(crate) fn try_new(list: Scalar, len: usize, ctx: &mut ExecutionCtx) -> VortexResult { - let data = PreparedSetData::try_new(list, ctx)?; + /// Prepares `elements`, the elements of a non-null list with the list nullability + /// `nullability`, as a set repeated `len` times. + pub(crate) fn try_new( + elements: ArrayRef, + nullability: Nullability, + len: usize, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let data = PreparedSetData::try_new(elements, nullability, ctx)?; Ok(Self::from_data(data, len)) } /// An array of `len` rows that share the prepared set `data`. fn from_data(data: PreparedSetData, len: usize) -> Self { - let dtype = data.list.dtype().clone(); + let dtype = data.list_dtype(); // SAFETY: the dtype is the dtype of the list that every row holds. unsafe { Array::from_parts_unchecked(ArrayParts::new(PreparedSet, dtype, len, data)) } @@ -226,47 +256,10 @@ impl Array { } const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ - ParentRuleSet::lift(&ConstantNeedleRule), ParentRuleSet::lift(&FilterReduceAdaptor(PreparedSet)), ParentRuleSet::lift(&SliceReduceAdaptor(PreparedSet)), ]); -/// Folds `list_contains` of a constant needle against the prepared set into its constant answer. -/// -/// [`ListContains`] folds a [`ConstantArray`] list the same way in its -/// [`reduce`](crate::scalar_fn::ScalarFnVTable::reduce). This rule covers the list after it is -/// prepared, which the generic constant check does not see. -#[derive(Debug)] -struct ConstantNeedleRule; - -impl ArrayParentReduceRule for ConstantNeedleRule { - type Parent = ExactScalarFn; - - fn reduce_parent( - &self, - array: ArrayView<'_, PreparedSet>, - parent: ScalarFnArrayView<'_, ListContains>, - child_idx: usize, - ) -> VortexResult> { - // The prepared set is the list child. As the needle it is a list of lists, which this rule - // does not fold. - if child_idx != 0 { - return Ok(None); - } - let scalar_fn_array = parent - .as_opt::() - .vortex_expect("ExactScalarFn matcher confirmed ScalarFnArray"); - let Some(needle) = scalar_fn_array.get_child(1).as_constant() else { - return Ok(None); - }; - - let result = compute_contains_scalar(array.list(), &needle, parent.options)?; - Ok(Some( - ConstantArray::new(result, scalar_fn_array.len()).into_array(), - )) - } -} - impl VTable for PreparedSet { type TypedArrayData = PreparedSetData; @@ -286,7 +279,7 @@ impl VTable for PreparedSet { _slots: &[Option], ) -> VortexResult<()> { vortex_ensure!( - data.list.dtype() == dtype, + &data.list_dtype() == dtype, "PreparedSetArray list dtype does not match outer dtype" ); Ok(()) @@ -337,10 +330,31 @@ impl VTable for PreparedSet { fn execute(array: Array, _ctx: &mut ExecutionCtx) -> VortexResult { // The rows hold the list only. The probe is of no use to the canonical form. - Ok(ExecutionResult::done(ConstantArray::new( - array.data().list.clone(), - array.len(), - ))) + let data = array.data(); + let n_elements = data.elements.len(); + + // Every row has the same offset and size, so use the narrowest width that fits the list. + let (offsets, sizes) = match_smallest_list_offset_type!(n_elements, |O| { + let size = + O::try_from(n_elements).vortex_expect("list length fits the chosen offset type"); + ( + ConstantArray::new::(O::default(), array.len()).into_array(), + ConstantArray::new::(size, array.len()).into_array(), + ) + }); + + // SAFETY: every view points at the range [0, n_elements) of the elements, and the list + // is never null. + let list = unsafe { + ListViewArray::new_unchecked( + data.elements.clone(), + offsets, + sizes, + Validity::from(data.nullability), + ) + }; + + Ok(ExecutionResult::done(list)) } fn reduce_parent( @@ -358,9 +372,17 @@ impl OperationsVTable for PreparedSet { fn scalar_at( array: ArrayView<'_, PreparedSet>, _index: usize, - _ctx: &mut ExecutionCtx, + ctx: &mut ExecutionCtx, ) -> VortexResult { - Ok(array.list.clone()) + let elements = (0..array.elements.len()) + .map(|idx| array.elements.execute_scalar(idx, ctx)) + .collect::>>()?; + + Ok(Scalar::list( + array.elements.dtype().clone(), + elements, + array.nullability, + )) } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs index 900c7ee0e0d..6a72dd5a88b 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -20,6 +20,7 @@ use vortex_buffer::BufferAllocatorRef; use vortex_buffer::BufferMut; use vortex_error::VortexExpect; use vortex_error::VortexResult; +use vortex_mask::Mask; use vortex_utils::aliases::hash_map::HashTable; use vortex_utils::aliases::hash_map::HashTableEntry; use vortex_utils::aliases::hash_map::RandomState; @@ -31,17 +32,17 @@ use crate::RecursiveCanonical; use crate::arrays::DecimalArray; use crate::arrays::PrimitiveArray; use crate::arrays::VarBinViewArray; -use crate::arrays::decimal::DecimalArrayExt; -use crate::arrays::decimal::widened_buffer; +use crate::arrays::decimal::converted_buffer; use crate::arrays::primitive::PrimitiveArrayExt; use crate::arrays::varbinview::BinaryView; use crate::dtype::DType; -use crate::dtype::IntegerPType; +use crate::dtype::DecimalType; +use crate::dtype::NativeDecimalType; +use crate::dtype::NativePType; use crate::dtype::PType; use crate::dtype::i256; use crate::match_each_decimal_value_type; use crate::match_each_integer_ptype; -use crate::scalar::DecimalValue; use crate::scalar_fn::fns::binary::build_row_comparator; use crate::scalar_fn::fns::binary::collect_bits; use crate::validity::Validity; @@ -52,187 +53,207 @@ const BITMAP_BITS_PER_ELEMENT: u128 = 64; /// A span this narrow is probed through a bitmap whatever the size of the set. const BITMAP_MIN_BITS: u128 = 1 << 12; -/// The structure that probes the non-null elements of a set. +/// Probes the non-null elements of a set for membership. /// -/// Every needle is probed against a bitmap, a sorted set or a hash table. Decimals use their -/// unscaled integers with the same bitmap and sorted-value strategy as primitives. Nested -/// values use sorted row indices with the same comparator as equality. No probe constructs -/// per-element expressions or materializes scalars in its loop. -enum Probe { - /// Integers, or floats by their bit patterns, spanning a dense range: one bit per value of the - /// span above `min_offset`, the smallest element as a `usize`. - Bitmap { - min_offset: usize, - span: usize, - bitmap: BitBuffer, - }, - /// Integers, or floats by their bit patterns, sorted without duplicates. - Sorted(PrimitiveArray), - /// Dense unscaled decimal values, including 128- and 256-bit storage. - DecimalBitmap { - min: DecimalValue, - bitmap: BitBuffer, - }, - /// Sorted, distinct unscaled decimal values. - DecimalSorted(DecimalArray), - /// UTF-8 or binary elements, found through a table of their indices hashed by their bytes, so - /// that no element is copied. - Bytes { - elements: VarBinViewArray, - hasher: RandomState, - table: HashTable, - }, - /// Recursively canonical elements, indexed in sorted order with duplicates removed. - Rows { - elements: ArrayRef, - indices: Vec, - }, -} - -impl Probe { - /// Builds the structure that probes the non-null `elements`. - fn try_new(elements: ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { - Ok(match elements.dtype() { - DType::Primitive(ptype, _) => { - let ptype = bit_pattern_ptype(*ptype); - let elements = elements - .execute::(ctx)? - .reinterpret_cast(ptype); - match_each_integer_ptype!(ptype, |T| { - integer_probe::(elements, ctx.allocator()) - }) - } - DType::Decimal(..) => { - let elements = elements.execute::(ctx)?; - match_each_decimal_value_type!(elements.values_type(), |T| { - let values = elements.buffer::(); - if let Some((min, bitmap)) = integer_bitmap(&values, ctx.allocator()) { - Probe::DecimalBitmap { - min: min.into(), - bitmap, - } - } else { - Probe::DecimalSorted(DecimalArray::new( - sorted_values(values), - elements.decimal_dtype(), - Validity::NonNullable, - )) - } - }) - } - DType::Utf8(_) | DType::Binary(_) => { - bytes_probe(elements.execute::(ctx)?) - } - _ => { - let elements = elements.execute::(ctx)?.0.into_array(); - let mut indices: Vec = (0..elements.len()).collect(); - if !indices.is_empty() { - let compare = build_row_comparator(&elements, &elements, ctx)?; - if !indices.is_sorted_by(|&lhs, &rhs| compare(lhs, rhs).is_le()) { - indices.sort_unstable_by(|&lhs, &rhs| compare(lhs, rhs)); - } - indices.dedup_by(|lhs, rhs| compare(*lhs, *rhs).is_eq()); - } - Probe::Rows { elements, indices } - } - }) - } - - /// One membership bit per needle, and the validity of the needles, which have the dtype of the - /// elements. +/// Integers, and floats by their bit patterns, are probed as an [`IntegerSet`] of their own type, +/// and decimals as an [`IntegerSet`] of their unscaled values. UTF-8 and binary values are found +/// through a hash table, and nested values through sorted row indices with the same comparator as +/// equality. No probe constructs per-element expressions or materializes scalars in its loop. +trait Probe: Send + Sync { + /// One membership bit per needle, and the validity of the needles, which have the dtype of + /// the elements. fn contains( &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, - ) -> VortexResult<(BitBuffer, Validity)> { + ) -> VortexResult<(BitBuffer, Validity)>; + + /// Whether the probe holds its values in a bitmap. + #[cfg(test)] + fn is_bitmap(&self) -> bool { + false + } +} + +/// Builds the probe of the non-null `elements`. +fn new_probe(elements: ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult> { + Ok(match elements.dtype() { + DType::Primitive(ptype, _) => { + let ptype = bit_pattern_ptype(*ptype); + let elements = elements + .execute::(ctx)? + .reinterpret_cast(ptype); + match_each_integer_ptype!(ptype, |T| { + let set = IntegerSet::::new(elements.into_buffer(), ctx.allocator()); + Box::new(PrimitiveSet(set)) as Box + }) + } + DType::Decimal(decimal, _) => { + let values_type = DecimalType::smallest_decimal_value_type(decimal); + let elements = elements.execute::(ctx)?; + let all_valid = Mask::new_true(elements.len()); + match_each_decimal_value_type!(values_type, |T| { + let values = converted_buffer::(&elements, &all_valid)?; + Box::new(DecimalSet(IntegerSet::new(values, ctx.allocator()))) as Box + }) + } + DType::Utf8(_) | DType::Binary(_) => { + Box::new(BytesSet::new(elements.execute::(ctx)?)) + } + _ => Box::new(RowSet::try_new(elements, ctx)?), + }) +} + +/// The distinct values of a set of integers of one type. +enum IntegerSet { + /// Values spanning a dense range: one bit per value of the span above `min`. + Bitmap { min: T, bitmap: BitBuffer }, + /// Values sorted without duplicates. + Sorted(Buffer), +} + +impl IntegerSet { + /// A bitmap over the values' span when the span is dense, and a sorted slice otherwise. + /// + /// A hash set and, for a handful of elements, a linear scan both lost to the binary search at + /// every set size measured by the `list_contains_set` benchmark, up to 16 384 elements. + fn new(values: Buffer, allocator: &BufferAllocatorRef) -> Self { + match integer_bitmap(&values, allocator) { + Some((min, bitmap)) => Self::Bitmap { min, bitmap }, + None => Self::Sorted(sorted_values(values)), + } + } + + /// One bit per needle, set when the needle is an element. + fn contains(&self, needles: &[T], allocator: &BufferAllocatorRef) -> BitBuffer { match self { - Probe::Bitmap { .. } | Probe::Sorted(_) => self.contains_primitive(needles, ctx), - Probe::DecimalBitmap { .. } | Probe::DecimalSorted(_) => { - self.contains_decimal(needles, ctx) - } - Probe::Bytes { - elements, - hasher, - table, - } => Self::contains_bytes(elements, hasher, table, needles, ctx), - Probe::Rows { elements, indices } => { - Self::contains_rows(elements, indices, needles, ctx) - } + Self::Bitmap { min, bitmap } => collect_bits( + needles, + // A needle below the smallest element wraps past the bitmap, so one comparison + // checks both bounds. + |needle| { + needle + .offset_from(*min) + .is_some_and(|offset| offset < bitmap.len() && bitmap.value(offset)) + }, + allocator, + ), + Self::Sorted(sorted) => collect_bits( + needles, + |needle| sorted.binary_search(&needle).is_ok(), + allocator, + ), } } +} + +/// Primitive integers, or floats by their bit patterns. +/// +/// A float is a member exactly when the compare kernel would call it equal to an element, which +/// is when their bit patterns match — distinguishing `-0.0` from `0.0` and one NaN payload from +/// another — so floats are probed by their bits, as integers. +struct PrimitiveSet(IntegerSet); - /// A float is a member exactly when the compare kernel would call it equal to an element, which - /// is when their bit patterns match — distinguishing `-0.0` from `0.0` and one NaN payload from - /// another — so floats are probed by their bits, as integers. - fn contains_primitive( +impl Probe for PrimitiveSet { + fn contains( &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, ) -> VortexResult<(BitBuffer, Validity)> { - let primitive = needles.clone().execute::(ctx)?; - let ptype = bit_pattern_ptype(primitive.ptype()); - let values = primitive.reinterpret_cast(ptype); - let bits = match_each_integer_ptype!(ptype, |T| { - self.integer_bits(values.as_slice::(), ctx.allocator()) - }); - Ok((bits, primitive.validity()?)) + let needles = needles + .clone() + .execute::(ctx)? + .reinterpret_cast(T::PTYPE); + let bits = self.0.contains(needles.as_slice::(), ctx.allocator()); + Ok((bits, needles.validity()?)) } - fn contains_decimal( + #[cfg(test)] + fn is_bitmap(&self) -> bool { + matches!(self.0, IntegerSet::Bitmap { .. }) + } +} + +/// The unscaled values of decimals, as the narrowest type that holds every value of their +/// precision. +/// +/// A needle has the precision of the elements, so a valid needle converts to that type whatever +/// its own storage width, and a needle stored at that type converts without a copy. +struct DecimalSet(IntegerSet); + +impl Probe for DecimalSet { + fn contains( &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, ) -> VortexResult<(BitBuffer, Validity)> { let needles = needles.clone().execute::(ctx)?; - // Logical precision and scale agree, but physical widths can differ. Widen once to - // the common width so out-of-range needles cannot truncate into false matches. - let bits = match self { - Probe::DecimalBitmap { min, bitmap } => { - let common = min.decimal_type().max(needles.values_type()); - match_each_decimal_value_type!(common, |T| { - let min = min.cast::().vortex_expect("lossless decimal widening"); - let values = widened_buffer::(&needles); - collect_bits( - &values, - |value| { - value - .offset_from(min) - .is_some_and(|offset| offset < bitmap.len() && bitmap.value(offset)) - }, - ctx.allocator(), - ) - }) - } - Probe::DecimalSorted(sorted) => { - let common = sorted.values_type().max(needles.values_type()); - match_each_decimal_value_type!(common, |T| { - let sorted = widened_buffer::(sorted); - let values = widened_buffer::(&needles); - sorted_bits(&sorted, &values, ctx.allocator()) - }) + let validity = needles.validity()?; + let valid = validity.execute_mask(needles.len(), ctx)?; + let values = converted_buffer::(&needles, &valid)?; + Ok((self.0.contains(&values, ctx.allocator()), validity)) + } + + #[cfg(test)] + fn is_bitmap(&self) -> bool { + matches!(self.0, IntegerSet::Bitmap { .. }) + } +} + +/// UTF-8 or binary elements, found through a table of their indices hashed by their bytes, so that +/// no element is copied. +struct BytesSet { + elements: VarBinViewArray, + hasher: RandomState, + table: HashTable, +} + +impl BytesSet { + /// A table of the elements' indices, hashed by their bytes, holding each distinct value once. + fn new(elements: VarBinViewArray) -> Self { + let hasher = RandomState::default(); + let mut table = HashTable::with_capacity(elements.len()); + { + let views = elements.views(); + let buffers = data_buffers(&elements); + let bytes = |idx: u32| view_bytes(&views[idx as usize], &buffers); + for (idx, view) in views.iter().enumerate() { + let value = view_bytes(view, &buffers); + if let HashTableEntry::Vacant(vacant) = table.entry( + hasher.hash_one(value), + |&other| bytes(other) == value, + |&other| hasher.hash_one(bytes(other)), + ) { + vacant.insert( + u32::try_from(idx).vortex_expect("a list holds fewer than 2^32 elements"), + ); + } } - _ => unreachable!("decimal needles meet a decimal probe"), - }; - Ok((bits, needles.validity()?)) + } + Self { + elements, + hasher, + table, + } } +} - fn contains_bytes( - elements: &VarBinViewArray, - hasher: &RandomState, - table: &HashTable, +impl Probe for BytesSet { + fn contains( + &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, ) -> VortexResult<(BitBuffer, Validity)> { - let element_views = elements.views(); - let element_buffers = data_buffers(elements); + let element_views = self.elements.views(); + let element_buffers = data_buffers(&self.elements); let array = needles.clone().execute::(ctx)?; let buffers = data_buffers(&array); let bits = collect_bits( array.views(), |view: BinaryView| { let value = view_bytes(&view, &buffers); - table - .find(hasher.hash_one(value), |&idx| { + self.table + .find(self.hasher.hash_one(value), |&idx| { view_bytes(&element_views[idx as usize], &element_buffers) == value }) .is_some() @@ -241,14 +262,36 @@ impl Probe { ); Ok((bits, array.validity()?)) } +} + +/// Recursively canonical elements, indexed in sorted order with duplicates removed. +struct RowSet { + elements: ArrayRef, + indices: Vec, +} + +impl RowSet { + fn try_new(elements: ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = elements.execute::(ctx)?.0.into_array(); + let mut indices: Vec = (0..elements.len()).collect(); + if !indices.is_empty() { + let compare = build_row_comparator(&elements, &elements, ctx)?; + if !indices.is_sorted_by(|&lhs, &rhs| compare(lhs, rhs).is_le()) { + indices.sort_unstable_by(|&lhs, &rhs| compare(lhs, rhs)); + } + indices.dedup_by(|lhs, rhs| compare(*lhs, *rhs).is_eq()); + } + Ok(Self { elements, indices }) + } +} - fn contains_rows( - elements: &ArrayRef, - indices: &[usize], +impl Probe for RowSet { + fn contains( + &self, needles: &ArrayRef, ctx: &mut ExecutionCtx, ) -> VortexResult<(BitBuffer, Validity)> { - if indices.is_empty() { + if self.indices.is_empty() { return Ok(( BitBuffer::full_in(false, needles.len(), ctx.allocator().clone()), needles.validity()?, @@ -261,11 +304,11 @@ impl Probe { .into_array(); // The comparator materializes child buffers and validity once, including decimal // widening and nested offsets. Each search then reads those buffers directly. - let compare = build_row_comparator(elements, &needles, ctx)?; + let compare = build_row_comparator(&self.elements, &needles, ctx)?; let bits = BitBuffer::collect_bool_in( needles.len(), |row| { - indices + self.indices .binary_search_by(|&element| compare(element, row)) .is_ok() }, @@ -273,68 +316,6 @@ impl Probe { ); Ok((bits, needles.validity()?)) } - - /// One bit per needle, set when the needle is an element. - fn integer_bits( - &self, - needles: &[T], - allocator: &BufferAllocatorRef, - ) -> BitBuffer { - match self { - Self::Bitmap { - min_offset, - span, - bitmap, - } => collect_bits( - needles, - // A needle below the smallest element wraps past `span`, so one comparison checks - // both bounds. - |needle| { - let offset = needle.as_().wrapping_sub(*min_offset); - offset <= *span && bitmap.value(offset) - }, - allocator, - ), - Self::Sorted(sorted) => { - let sorted = sorted.as_slice::(); - sorted_bits(sorted, needles, allocator) - } - Self::DecimalBitmap { .. } - | Self::DecimalSorted(_) - | Self::Bytes { .. } - | Self::Rows { .. } => { - unreachable!("integer needles meet an integer probe") - } - } - } -} - -/// A table of the elements' indices, hashed by their bytes, holding each distinct value once. -fn bytes_probe(elements: VarBinViewArray) -> Probe { - let hasher = RandomState::default(); - let mut table = HashTable::with_capacity(elements.len()); - { - let views = elements.views(); - let buffers = data_buffers(&elements); - let bytes = |idx: u32| view_bytes(&views[idx as usize], &buffers); - for (idx, view) in views.iter().enumerate() { - let value = view_bytes(view, &buffers); - if let HashTableEntry::Vacant(vacant) = table.entry( - hasher.hash_one(value), - |&other| bytes(other) == value, - |&other| hasher.hash_one(bytes(other)), - ) { - vacant.insert( - u32::try_from(idx).vortex_expect("a list holds fewer than 2^32 elements"), - ); - } - } - } - Probe::Bytes { - elements, - hasher, - table, - } } /// The host slices of an array's data buffers, indexed by a view's buffer index. @@ -364,33 +345,9 @@ fn bit_pattern_ptype(ptype: PType) -> PType { } } -/// A bitmap over the elements' span when the span is dense, and a sorted slice otherwise. -/// -/// A hash set and, for a handful of elements, a linear scan both lost to the binary search at every -/// set size measured by the `list_contains_set` benchmark, up to 16 384 elements. -fn integer_probe( - elements: PrimitiveArray, - allocator: &BufferAllocatorRef, -) -> Probe { - // The primitive hot loop computes offsets in usize, which must hold every value of T. - if size_of::() <= size_of::() - && let Some((min, bitmap)) = integer_bitmap(elements.as_slice::(), allocator) - { - return Probe::Bitmap { - min_offset: min.as_(), - span: bitmap.len() - 1, - bitmap, - }; - } - Probe::Sorted(PrimitiveArray::new( - sorted_values(elements.into_buffer::()), - Validity::NonNullable, - )) -} - /// An unsigned modular distance, rejecting offsets too wide to address a bitmap. /// Keeping the subtraction at the physical width also handles signed ranges spanning zero. -trait SetInteger: Copy + Ord { +trait SetInteger: Copy + Ord + Send + Sync + 'static { fn offset_from(self, min: Self) -> Option; } @@ -452,17 +409,5 @@ fn sorted_values(values: Buffer) -> Bu Buffer::from(sorted) } -fn sorted_bits( - sorted: &[T], - needles: &[T], - allocator: &BufferAllocatorRef, -) -> BitBuffer { - collect_bits( - needles, - |needle| sorted.binary_search(&needle).is_ok(), - allocator, - ) -} - #[cfg(test)] mod tests; diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs index c08511979a3..788e478a612 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -8,8 +8,8 @@ use vortex_error::VortexResult; use super::PreparedSet; use super::PreparedSetArray; -use super::Probe; use crate::ArrayRef; +use crate::ExecutionCtx; use crate::IntoArray; use crate::VortexSessionExecute; use crate::array_session; @@ -21,6 +21,7 @@ use crate::arrays::FixedSizeListArray; use crate::arrays::ListArray; use crate::arrays::PrimitiveArray; use crate::arrays::StructArray; +use crate::arrays::constant::list_scalar_elements; use crate::assert_arrays_eq; use crate::builders::builder_with_capacity_in; use crate::dtype::DType; @@ -30,7 +31,6 @@ use crate::dtype::Nullability; use crate::dtype::PType; use crate::dtype::i256; use crate::match_each_decimal_value_type; -use crate::optimizer::ArrayOptimizer; use crate::scalar::DecimalValue; use crate::scalar::Scalar; use crate::scalar_fn::fns::list_contains::ListContains; @@ -38,6 +38,18 @@ use crate::scalar_fn::fns::list_contains::ListContainsOptions; use crate::scalar_fn::fns::list_contains::PreparedSetData; use crate::validity::Validity; +/// Prepares the elements of the non-null list scalar `list`. +fn prepare(list: &Scalar, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + PreparedSetData::try_new(elements, list.dtype().nullability(), ctx) +} + +/// Prepares the elements of the non-null list scalar `list` as a set repeated `len` times. +fn prepare_array(list: &Scalar, len: usize, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + Ok(PreparedSetArray::try_new(elements, list.dtype().nullability(), len, ctx)?.into_array()) +} + fn nested_needles() -> ArrayRef { ListArray::try_new( PrimitiveArray::from_option_iter([Some(1i32), None, Some(2), None, Some(9)]).into_array(), @@ -103,7 +115,7 @@ fn test_row_set_returns_membership_bits( elements.push(Scalar::null(dtype.clone())); let list = Scalar::list(dtype, elements, Nullability::NonNullable); let options = ListContainsOptions { sql_null_semantics }; - let set = PreparedSetData::try_new(list, &mut ctx)?; + let set = prepare(&list, &mut ctx)?; let result = set.contains(&needles, &options, &mut ctx)?; // The old fallback returned a lazy OR tree here. assert!(result.is::()); @@ -170,8 +182,8 @@ fn test_decimal_bitmap_across_storage_widths( .into_array() }); let options = ListContainsOptions { sql_null_semantics }; - let set = PreparedSetData::try_new(list, &mut ctx)?; - assert!(matches!(&set.set.probe, Probe::DecimalBitmap { .. })); + let set = prepare(&list, &mut ctx)?; + assert!(set.set.probe.is_bitmap()); let non_match = (!sql_null_semantics).then_some(false); assert_arrays_eq!( set.contains(&needles, &options, &mut ctx)?, @@ -214,9 +226,8 @@ fn test_decimal_wide_values( decimal, ) .into_array(); - let set = PreparedSetData::try_new(list, &mut ctx)?; - assert_eq!(matches!(&set.set.probe, Probe::DecimalBitmap { .. }), dense); - assert_eq!(matches!(&set.set.probe, Probe::DecimalSorted(_)), !dense); + let set = prepare(&list, &mut ctx)?; + assert_eq!(set.set.probe.is_bitmap(), dense); assert_arrays_eq!( set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, BoolArray::from_iter([ @@ -232,6 +243,44 @@ fn test_decimal_wide_values( Ok(()) } +#[rstest] +#[case::bitmap( + PrimitiveArray::from_option_iter([Some(5i64), Some(-3), Some(5)]).into_array(), + PrimitiveArray::from_option_iter([Some(-3i64), Some(4), Some(i64::MIN), None]).into_array(), + true, + [Some(true), Some(false), Some(false), None], +)] +#[case::unsigned_in_signed_order( + PrimitiveArray::from_option_iter([Some(u64::MAX), Some(0), Some(1 << 40)]).into_array(), + PrimitiveArray::from_option_iter([Some(1u64 << 40), Some(1), Some(u64::MAX), Some(0)]) + .into_array(), + false, + [Some(true), Some(false), Some(true), Some(true)], +)] +#[case::float_bits( + PrimitiveArray::from_option_iter([Some(0.0f64), Some(f64::NAN)]).into_array(), + PrimitiveArray::from_option_iter([Some(-0.0f64), Some(0.0), Some(f64::NAN), Some(1.5)]) + .into_array(), + false, + [Some(false), Some(true), Some(true), Some(false)], +)] +fn test_primitive_integers( + #[case] elements: ArrayRef, + #[case] needles: ArrayRef, + #[case] dense: bool, + #[case] expected: [Option; 4], +) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let set = PreparedSetData::try_new(elements, Nullability::NonNullable, &mut ctx)?; + assert_eq!(set.set.probe.is_bitmap(), dense); + assert_arrays_eq!( + set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, + BoolArray::from_iter(expected), + &mut ctx + ); + Ok(()) +} + /// The set `{2, null}` of nullable `i32`. fn set_with_null() -> Scalar { let element = DType::Primitive(PType::I32, Nullability::Nullable); @@ -248,7 +297,7 @@ fn set_with_null() -> Scalar { #[test] fn test_prepared_set_rows_are_the_constant_list() -> VortexResult<()> { let mut ctx = array_session().create_execution_ctx(); - let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + let set = prepare_array(&set_with_null(), 4, &mut ctx)?; assert_arrays_eq!(set, ConstantArray::new(set_with_null(), 4), &mut ctx); @@ -270,7 +319,7 @@ fn test_list_contains_probes_a_prepared_set_list( #[case] expected: [Option; 4], ) -> VortexResult<()> { let mut ctx = array_session().create_execution_ctx(); - let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + let set = prepare_array(&set_with_null(), 4, &mut ctx)?; let needles = PrimitiveArray::from_option_iter([Some(1i32), Some(2), None, Some(3)]).into_array(); @@ -282,23 +331,43 @@ fn test_list_contains_probes_a_prepared_set_list( #[rstest] #[case::default(ListContainsOptions::default(), Some(false))] #[case::sql(ListContainsOptions { sql_null_semantics: true }, None)] -fn test_constant_needle_folds_without_probing( +fn test_constant_needle_gives_a_constant( #[case] options: ListContainsOptions, #[case] expected: Option, ) -> VortexResult<()> { - // `3 IN (2, NULL)` against the prepared set is decided at optimization, as it is against a - // constant list. + // `3 IN (2, NULL)` against the prepared set is looked up once, and gives a constant. let mut ctx = array_session().create_execution_ctx(); - let set = PreparedSetArray::try_new(set_with_null(), 4, &mut ctx)?.into_array(); + let set = prepare_array(&set_with_null(), 4, &mut ctx)?; let needle = ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 4).into_array(); - let optimized = ListContains::try_new_opts(set, needle, options)? + let result = ListContains::try_new_opts(set, needle, options)? .into_array() - .optimize()?; + .execute::(&mut ctx)?; let expected = match expected { Some(value) => Scalar::bool(value, Nullability::Nullable), None => Scalar::null(DType::Bool(Nullability::Nullable)), }; - assert_eq!(optimized.as_constant(), Some(expected)); + assert_eq!(result.as_constant(), Some(expected)); + Ok(()) +} + +#[test] +fn test_constant_row_needle_probes_one_row() -> VortexResult<()> { + // A probe of rows cannot look a scalar needle up, so execution probes one row. + let mut ctx = array_session().create_execution_ctx(); + let needles = nested_needles(); + let dtype = needles.dtype().as_nullable(); + let member = needles.execute_scalar(2, &mut ctx)?.cast(&dtype)?; + let list = Scalar::list(dtype, vec![member.clone()], Nullability::NonNullable); + + let set = prepare_array(&list, 4, &mut ctx)?; + let needle = ConstantArray::new(member, 4).into_array(); + let result = ListContains::try_new_opts(set, needle, ListContainsOptions::default())? + .into_array() + .execute::(&mut ctx)?; + assert_eq!( + result.as_constant(), + Some(Scalar::bool(true, Nullability::Nullable)) + ); Ok(()) } From b566df28889e8076fea947589b534cdb043e391a Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 22:51:00 -0400 Subject: [PATCH 6/9] less Signed-off-by: Robert Kruszewski --- .../src/arrays/constant/vtable/canonical.rs | 64 ++------------- .../src/scalar_fn/fns/list_contains/kernel.rs | 7 +- .../fns/list_contains/prepared/array.rs | 77 +++++++++++++++---- .../fns/list_contains/prepared/tests.rs | 41 ++++++++++ 4 files changed, 114 insertions(+), 75 deletions(-) diff --git a/vortex-array/src/arrays/constant/vtable/canonical.rs b/vortex-array/src/arrays/constant/vtable/canonical.rs index c7241a9c753..409224d1b86 100644 --- a/vortex-array/src/arrays/constant/vtable/canonical.rs +++ b/vortex-array/src/arrays/constant/vtable/canonical.rs @@ -7,7 +7,6 @@ use itertools::Itertools; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; use vortex_buffer::BufferAllocatorRef; -use vortex_buffer::BufferMut; use vortex_buffer::BufferString; use vortex_buffer::ByteBuffer; use vortex_buffer::buffer; @@ -35,8 +34,6 @@ use crate::arrays::UnionArray; use crate::arrays::VarBinViewArray; use crate::arrays::VariantArray; use crate::arrays::varbinview::BinaryView; -use crate::builders::ArrayBuilder; -use crate::builders::VarBinViewBuilder; use crate::builders::builder_with_capacity_in; use crate::dtype::DType; use crate::dtype::DecimalType; @@ -48,7 +45,6 @@ use crate::match_smallest_list_offset_type; use crate::scalar::DecimalValue; use crate::scalar::ListScalar; use crate::scalar::Scalar; -use crate::scalar::ScalarValue; use crate::validity::Validity; /// Shared implementation for both `canonicalize` and `execute` methods. @@ -310,65 +306,19 @@ fn constant_canonical_list_array( } /// The elements of a list scalar as an array, one row per element; empty for a null list. -/// -/// Primitive and byte elements are written straight from their values rather than through a scalar -/// each. pub(crate) fn list_scalar_elements(list: &ListScalar, allocator: &BufferAllocatorRef) -> ArrayRef { let element_dtype = list.element_dtype(); - let Some(values) = list.element_values() else { + let Some(elements) = list.elements() else { return Canonical::empty(element_dtype).into_array(); }; - match element_dtype { - DType::Primitive(ptype, nullability) => match_each_native_ptype!(ptype, |T| { - let mut buffer = BufferMut::::with_capacity_in(values.len(), allocator.clone()); - buffer.extend(values.iter().map(|value| { - value.as_ref().map_or_else(T::default, |value| { - value - .as_primitive() - .cast::() - .vortex_expect("list element of the list's element ptype") - }) - })); - PrimitiveArray::new(buffer.freeze(), element_validity(values, *nullability)) - .into_array() - }), - DType::Utf8(_) | DType::Binary(_) => { - let mut builder = VarBinViewBuilder::with_capacity_in( - element_dtype.clone(), - values.len(), - allocator.clone(), - ); - for value in values { - match value { - None => builder.append_null(), - Some(ScalarValue::Utf8(value)) => builder.append_value(value.as_bytes()), - Some(value) => builder.append_value(value.as_binary().as_slice()), - } - } - builder.finish() - } - _ => { - let mut builder = builder_with_capacity_in(element_dtype, values.len(), allocator); - for idx in 0..values.len() { - builder - .append_scalar(&list.element(idx).vortex_expect("index within the list")) - .vortex_expect("list element scalar was invalid"); - } - builder.finish() - } - } -} - -/// The validity of a list's elements, a null element being a `None` value. -fn element_validity(values: &[Option], nullability: Nullability) -> Validity { - match nullability { - Nullability::NonNullable => Validity::NonNullable, - Nullability::Nullable if values.iter().all(Option::is_some) => Validity::AllValid, - Nullability::Nullable => { - Validity::from(BitBuffer::from_iter(values.iter().map(Option::is_some))) - } + let mut builder = builder_with_capacity_in(element_dtype, elements.len(), allocator); + for element in &elements { + builder + .append_scalar(element) + .vortex_expect("list element scalar was invalid"); } + builder.finish() } /// Creates a [`FixedSizeListArray`] whose every row holds the same list. 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 b4ba2ee520b..952ac16d8c9 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -41,10 +41,15 @@ pub trait ListContainsElementReduce: VTable { /// For a needle that is not canonical, execution prepares a constant list into a /// [`PreparedSetArray`] and runs the kernels again. Thus a kernel can get the prepared set with /// `list.as_opt::()`, and probe its own values with [`PreparedSetData::contains`]. -/// For example, a dictionary probes only its values. +/// For example, a dictionary probes only its values. A single value, such as the fill value of a +/// sparse needle, is probed with [`PreparedSetData::contains_scalar`]. A kernel that finds its +/// matches in another way makes its result with [`PreparedSetData::result_from_bits`], which applies +/// the same null semantics. /// /// [`PreparedSetArray`]: crate::scalar_fn::fns::list_contains::PreparedSetArray /// [`PreparedSetData::contains`]: crate::scalar_fn::fns::list_contains::PreparedSetData::contains +/// [`PreparedSetData::contains_scalar`]: crate::scalar_fn::fns::list_contains::PreparedSetData::contains_scalar +/// [`PreparedSetData::result_from_bits`]: crate::scalar_fn::fns::list_contains::PreparedSetData::result_from_bits pub trait ListContainsElementKernel: VTable { fn list_contains( list: &ArrayRef, diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs index d276fbe10de..d30e2a09380 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -142,54 +142,74 @@ impl PreparedSetData { /// a null element gives `null` for a needle that matches no element. /// /// A constant needle is probed once, and gives a constant result. + /// + /// # Errors + /// + /// Fails when the needles do not have the dtype of the list's elements. pub fn contains( &self, needles: &ArrayRef, options: &ListContainsOptions, ctx: &mut ExecutionCtx, ) -> VortexResult { - let element_dtype = self.elements.dtype(); - if !element_dtype.eq_ignore_nullability(needles.dtype()) { - vortex_bail!( - "Element type {} of list does not match search value {}", - element_dtype, - needles.dtype(), - ); - } - if let Some(needle) = needles.as_constant() { let result = self.contains_scalar(&needle, options, ctx)?; return Ok(ConstantArray::new(result, needles.len()).into_array()); } + self.check_needle_dtype(needles.dtype())?; let (bits, needle_validity) = self.set.probe.contains(needles, ctx)?; - self.finish(bits, needle_validity, needles.dtype(), options) + self.result_from_bits(bits, needle_validity, needles.dtype(), options) } /// Whether the constant `needle` is an element of the list, under `options`. /// - /// The answer is the one [`Self::contains`] gives for each row of a needle with this value. The - /// probe gets one row, not a row per needle. - fn contains_scalar( + /// The answer is the one [`Self::contains`] gives for each row of a needle with this value, for + /// example for the fill value of a sparse needle. The probe gets one row. + /// + /// # Errors + /// + /// Fails when the needle does not have the dtype of the list's elements. + pub fn contains_scalar( &self, needle: &Scalar, options: &ListContainsOptions, ctx: &mut ExecutionCtx, ) -> VortexResult { + self.check_needle_dtype(needle.dtype())?; let needles = ConstantArray::new(needle.clone(), 1).into_array(); let (bits, needle_validity) = self.set.probe.contains(&needles, ctx)?; - self.finish(bits, needle_validity, needle.dtype(), options)? + self.result_from_bits(bits, needle_validity, needle.dtype(), options)? .execute_scalar(0, ctx) } - /// Makes the result from one membership bit per needle and the validity of the needles. - fn finish( + /// Makes the result of [`Self::contains`] from one membership bit per needle. + /// + /// A kernel that finds the matches of its needles without this probe, for example from bounds + /// of its own, uses this to apply the null semantics of [`Self::contains`]. A set bit in `bits` + /// tells that the needle equals an element. `needle_validity` is the validity of the needles, + /// and `needle_dtype` is their dtype. The bit of a null needle has no effect. + /// + /// # Errors + /// + /// Fails when `needle_dtype` is not the dtype of the list's elements, or when + /// `needle_validity` does not have one row per bit. + pub fn result_from_bits( &self, bits: BitBuffer, needle_validity: Validity, needle_dtype: &DType, options: &ListContainsOptions, ) -> VortexResult { + self.check_needle_dtype(needle_dtype)?; + if let Some(len) = needle_validity.maybe_len() { + vortex_ensure!( + len == bits.len(), + "Needle validity has {len} rows, but there are {} membership bits", + bits.len() + ); + } + let nullability = options.result_nullability(&self.list_dtype(), needle_dtype); let validity = if self.set.is_empty && !options.sql_null_semantics { @@ -203,6 +223,19 @@ impl PreparedSetData { Ok(BoolArray::new(bits, validity.union_nullability(nullability)).into_array()) } + + /// Fails when needles of `needle_dtype` cannot be elements of the list. + fn check_needle_dtype(&self, needle_dtype: &DType) -> VortexResult<()> { + let element_dtype = self.elements.dtype(); + if !element_dtype.eq_ignore_nullability(needle_dtype) { + vortex_bail!( + "Element type {} of list does not match search value {}", + element_dtype, + needle_dtype, + ); + } + Ok(()) + } } impl Debug for PreparedSetData { @@ -236,7 +269,17 @@ impl ArrayEq for PreparedSetData { impl Array { /// Prepares `elements`, the elements of a non-null list with the list nullability /// `nullability`, as a set repeated `len` times. - pub(crate) fn try_new( + /// + /// [`ListContains`] prepares a constant list itself. A kernel crate can use this to test its + /// [`ListContainsElementKernel`] against a prepared set. + /// + /// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains + /// [`ListContainsElementKernel`]: crate::scalar_fn::fns::list_contains::ListContainsElementKernel + /// + /// # Errors + /// + /// Fails when the elements cannot be executed into the probe. + pub fn try_new( elements: ArrayRef, nullability: Nullability, len: usize, diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs index 788e478a612..832580ccb17 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -371,3 +371,44 @@ fn test_constant_row_needle_probes_one_row() -> VortexResult<()> { ); Ok(()) } + +#[rstest] +#[case::default(ListContainsOptions::default(), [Some(true), Some(false), None])] +#[case::sql(ListContainsOptions { sql_null_semantics: true }, [Some(true), None, None])] +fn test_result_from_bits_applies_null_semantics( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 3], +) -> VortexResult<()> { + // Against `{2, null}`: a match, a non-match, and a null needle whose bit has no effect. + let mut ctx = array_session().create_execution_ctx(); + let set = prepare(&set_with_null(), &mut ctx)?; + let needle_dtype = DType::Primitive(PType::I32, Nullability::Nullable); + let bits = BitBuffer::from_iter([true, false, true]); + let validity = Validity::from_iter([true, true, false]); + + let result = set.result_from_bits(bits, validity, &needle_dtype, &options)?; + assert_arrays_eq!(result, BoolArray::from_iter(expected), &mut ctx); + Ok(()) +} + +#[test] +fn test_result_from_bits_rejects_mismatched_needles() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let set = prepare(&set_with_null(), &mut ctx)?; + let options = ListContainsOptions::default(); + let bits = BitBuffer::from_iter([true, false]); + + let short_validity = Validity::from_iter([true]); + let needle_dtype = DType::Primitive(PType::I32, Nullability::Nullable); + assert!( + set.result_from_bits(bits.clone(), short_validity, &needle_dtype, &options) + .is_err() + ); + + let wrong_dtype = DType::Primitive(PType::I64, Nullability::Nullable); + assert!( + set.result_from_bits(bits, Validity::AllValid, &wrong_dtype, &options) + .is_err() + ); + Ok(()) +} From 159aaff54496c65ea6596c98c2cf4c2aae49d840 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Tue, 29 Sep 2026 12:12:56 -0700 Subject: [PATCH 7/9] chunkedcompute Signed-off-by: Robert Kruszewski --- .../src/arrays/chunked/compute/kernel.rs | 7 + .../arrays/chunked/compute/list_contains.rs | 129 +++++++++++++ .../src/arrays/chunked/compute/mod.rs | 1 + .../src/arrays/chunked/compute/rules.rs | 9 + .../src/arrays/constant/vtable/canonical.rs | 19 +- .../fns/list_contains/prepared/mod.rs | 180 ++++++++++++++---- .../fns/list_contains/prepared/tests.rs | 44 +++++ 7 files changed, 345 insertions(+), 44 deletions(-) create mode 100644 vortex-array/src/arrays/chunked/compute/list_contains.rs diff --git a/vortex-array/src/arrays/chunked/compute/kernel.rs b/vortex-array/src/arrays/chunked/compute/kernel.rs index db0042105cd..e8f1986c6d2 100644 --- a/vortex-array/src/arrays/chunked/compute/kernel.rs +++ b/vortex-array/src/arrays/chunked/compute/kernel.rs @@ -13,6 +13,8 @@ use crate::arrays::filter::FilterExecuteAdaptor; use crate::arrays::slice::SliceExecuteAdaptor; use crate::optimizer::kernels::ArrayKernelsExt; use crate::scalar_fn::ScalarFnVTable; +use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsElementExecuteAdaptor; use crate::scalar_fn::fns::mask::Mask; use crate::scalar_fn::fns::mask::MaskExecuteAdaptor; use crate::scalar_fn::fns::zip::Zip; @@ -25,4 +27,9 @@ pub(crate) fn initialize(session: &VortexSession) { kernels.register_execute_parent_kernel(Slice.id(), Chunked, SliceExecuteAdaptor(Chunked)); kernels.register_execute_parent_kernel(Dict.id(), Chunked, TakeExecuteAdaptor(Chunked)); kernels.register_execute_parent_kernel(Zip.id(), Chunked, ZipExecuteAdaptor(Chunked)); + kernels.register_execute_parent_kernel( + ListContains.id(), + Chunked, + ListContainsElementExecuteAdaptor(Chunked), + ); } diff --git a/vortex-array/src/arrays/chunked/compute/list_contains.rs b/vortex-array/src/arrays/chunked/compute/list_contains.rs new file mode 100644 index 00000000000..03840fef761 --- /dev/null +++ b/vortex-array/src/arrays/chunked/compute/list_contains.rs @@ -0,0 +1,129 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use vortex_error::VortexResult; + +use crate::ArrayRef; +use crate::ExecutionCtx; +use crate::IntoArray; +use crate::array::ArrayView; +use crate::arrays::Chunked; +use crate::arrays::ChunkedArray; +use crate::arrays::chunked::ChunkedArrayExt; +use crate::dtype::DType; +use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsElementKernel; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::scalar_fn::fns::list_contains::PreparedSet; + +/// Probes each chunk of the needles against one prepared set. +/// +/// A constant list is prepared first, by the execution of [`ListContains`]. Each chunk then gets a +/// lazy [`ListContains`] over a slice of the prepared set, which shares the one probe, so that the +/// kernels of the chunk's own encoding can probe it. +impl ListContainsElementKernel for Chunked { + fn list_contains( + list: &ArrayRef, + needles: ArrayView<'_, Chunked>, + options: &ListContainsOptions, + _ctx: &mut ExecutionCtx, + ) -> VortexResult> { + if !list.is::() { + return Ok(None); + } + + let mut offset = 0; + let chunks = needles + .iter_chunks() + .map(|chunk| { + let set = list.slice(offset..offset + chunk.len())?; + offset += chunk.len(); + Ok(ListContains::try_new_opts(set, chunk.clone(), *options)?.into_array()) + }) + .collect::>>()?; + + let dtype = DType::Bool(options.result_nullability(list.dtype(), needles.dtype())); + + // SAFETY: each chunk is `list_contains` of the prepared set and a needle chunk of one dtype, + // so every chunk has the dtype of the whole result. + Ok(Some( + unsafe { ChunkedArray::new_unchecked(chunks, dtype) }.into_array(), + )) + } +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + use vortex_error::VortexResult; + + use crate::ArrayRef; + use crate::IntoArray; + use crate::VortexSessionExecute; + use crate::array_session; + use crate::arrays::BoolArray; + use crate::arrays::Chunked; + use crate::arrays::ChunkedArray; + use crate::arrays::ConstantArray; + use crate::arrays::PrimitiveArray; + use crate::assert_arrays_eq; + use crate::dtype::DType; + use crate::dtype::Nullability; + use crate::dtype::PType; + use crate::optimizer::ArrayOptimizer; + use crate::scalar::Scalar; + use crate::scalar_fn::fns::list_contains::ListContains; + use crate::scalar_fn::fns::list_contains::ListContainsOptions; + + /// The set `{2, null}` of nullable `i32`. + fn set_with_null() -> Scalar { + let element = DType::Primitive(PType::I32, Nullability::Nullable); + Scalar::list( + element.clone(), + vec![ + Scalar::primitive(2i32, Nullability::Nullable), + Scalar::null(element), + ], + Nullability::NonNullable, + ) + } + + fn chunked_needles() -> VortexResult { + Ok(ChunkedArray::try_new( + vec![ + PrimitiveArray::from_option_iter([Some(1i32), Some(2)]).into_array(), + PrimitiveArray::from_option_iter::([]).into_array(), + PrimitiveArray::from_option_iter([None, Some(3), Some(2)]).into_array(), + ], + DType::Primitive(PType::I32, Nullability::Nullable), + )? + .into_array()) + } + + #[rstest] + #[case::default( + ListContainsOptions::default(), + [Some(false), Some(true), None, Some(false), Some(true)], + )] + #[case::sql( + ListContainsOptions { sql_null_semantics: true }, + [None, Some(true), None, None, Some(true)], + )] + fn test_constant_list_over_chunked_needles( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 5], + ) -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let needles = chunked_needles()?; + let list = ConstantArray::new(set_with_null(), needles.len()).into_array(); + + // The node stays whole, so that execution prepares the list once for every chunk. + let array = ListContains::try_new_opts(list, needles, options)? + .into_array() + .optimize()?; + assert!(!array.is::()); + + assert_arrays_eq!(array, BoolArray::from_iter(expected), &mut ctx); + Ok(()) + } +} diff --git a/vortex-array/src/arrays/chunked/compute/mod.rs b/vortex-array/src/arrays/chunked/compute/mod.rs index 66d64595814..7e75109b1f1 100644 --- a/vortex-array/src/arrays/chunked/compute/mod.rs +++ b/vortex-array/src/arrays/chunked/compute/mod.rs @@ -6,6 +6,7 @@ mod cast; mod fill_null; mod filter; pub(crate) mod kernel; +mod list_contains; mod mask; pub(crate) mod rules; mod slice; diff --git a/vortex-array/src/arrays/chunked/compute/rules.rs b/vortex-array/src/arrays/chunked/compute/rules.rs index 712578255bb..f92cdeed236 100644 --- a/vortex-array/src/arrays/chunked/compute/rules.rs +++ b/vortex-array/src/arrays/chunked/compute/rules.rs @@ -21,6 +21,7 @@ use crate::optimizer::rules::ArrayParentReduceRule; use crate::optimizer::rules::ParentRuleSet; use crate::scalar_fn::fns::cast::CastReduceAdaptor; use crate::scalar_fn::fns::fill_null::FillNullReduceAdaptor; +use crate::scalar_fn::fns::list_contains::ListContains; pub(crate) const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ ParentRuleSet::lift(&CastReduceAdaptor(Chunked)), @@ -61,6 +62,10 @@ impl ArrayParentReduceRule for ChunkedUnaryScalarFnPushDownRule { } /// Push down non-unary scalar functions through chunked arrays where other siblings are constant. +/// +/// [`ListContains`] is not pushed down. Its execution prepares a constant list once as a set, and +/// its chunked kernel then probes each chunk against that one set. A push-down would give each +/// chunk its own copy of the list, and so its own set to prepare. #[derive(Debug)] struct ChunkedConstantScalarFnPushDownRule; impl ArrayParentReduceRule for ChunkedConstantScalarFnPushDownRule { @@ -72,6 +77,10 @@ impl ArrayParentReduceRule for ChunkedConstantScalarFnPushDownRule { parent: ArrayView<'_, ScalarFn>, child_idx: usize, ) -> VortexResult> { + if parent.scalar_fn().is::() { + return Ok(None); + } + for (idx, child) in parent.iter_children().enumerate() { if idx == child_idx { continue; diff --git a/vortex-array/src/arrays/constant/vtable/canonical.rs b/vortex-array/src/arrays/constant/vtable/canonical.rs index 409224d1b86..7c5ff18106f 100644 --- a/vortex-array/src/arrays/constant/vtable/canonical.rs +++ b/vortex-array/src/arrays/constant/vtable/canonical.rs @@ -308,15 +308,24 @@ fn constant_canonical_list_array( /// The elements of a list scalar as an array, one row per element; empty for a null list. pub(crate) fn list_scalar_elements(list: &ListScalar, allocator: &BufferAllocatorRef) -> ArrayRef { let element_dtype = list.element_dtype(); - let Some(elements) = list.elements() else { + let Some(elements) = list.element_values() else { return Canonical::empty(element_dtype).into_array(); }; let mut builder = builder_with_capacity_in(element_dtype, elements.len(), allocator); - for element in &elements { - builder - .append_scalar(element) - .vortex_expect("list element scalar was invalid"); + for element in elements { + match element { + Some(element) => { + builder + .append_scalar(&unsafe { + Scalar::new_unchecked(element_dtype.clone(), Some(element.clone())) + }) + .vortex_expect("list element scalar was invalid"); + } + None => { + builder.append_null(); + } + } } builder.finish() } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs index 6a72dd5a88b..0d3028c8e01 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -52,6 +52,9 @@ use crate::validity::Validity; const BITMAP_BITS_PER_ELEMENT: u128 = 64; /// A span this narrow is probed through a bitmap whatever the size of the set. const BITMAP_MIN_BITS: u128 = 1 << 12; +/// The bits per element of the filter of view heads. With one bit per head, about one in eight +/// needles that are not elements passes the filter of a large set. +const HEAD_FILTER_BITS_PER_ELEMENT: usize = 8; /// Probes the non-null elements of a set for membership. /// @@ -200,42 +203,107 @@ impl Probe for DecimalSet { } } -/// UTF-8 or binary elements, found through a table of their indices hashed by their bytes, so that -/// no element is copied. +/// UTF-8 or binary elements, found by their views, so that most needles are decided without a +/// read of a data buffer or a hash of their bytes. +/// +/// The first 8 bytes of a view, its head, hold the length of the value and its first 4 bytes, +/// zero-padded, whether the value is inlined or not. A value of at most 12 bytes is inlined whole, +/// zero-padded as the compare kernel requires. Thus a needle is probed in three steps, the +/// cheapest first: +/// +/// 1. The filter of the elements' heads rejects most non-members with one multiply and one bit. +/// 2. A short needle is an element exactly when its whole view is the view of a short element. +/// 3. A long needle is found through a table of the long elements, hashed by their bytes. Its +/// head is compared before the bytes after the prefix. struct BytesSet { - elements: VarBinViewArray, + heads: HeadFilter, hasher: RandomState, - table: HashTable, + /// The distinct whole views of the elements of at most 12 bytes. + short: HashTable, + /// The elements, for the bytes of the long ones. + elements: VarBinViewArray, + /// The indices of the distinct long elements, hashed by their bytes. + long: HashTable, } impl BytesSet { - /// A table of the elements' indices, hashed by their bytes, holding each distinct value once. fn new(elements: VarBinViewArray) -> Self { let hasher = RandomState::default(); - let mut table = HashTable::with_capacity(elements.len()); - { - let views = elements.views(); - let buffers = data_buffers(&elements); - let bytes = |idx: u32| view_bytes(&views[idx as usize], &buffers); - for (idx, view) in views.iter().enumerate() { - let value = view_bytes(view, &buffers); - if let HashTableEntry::Vacant(vacant) = table.entry( - hasher.hash_one(value), - |&other| bytes(other) == value, - |&other| hasher.hash_one(bytes(other)), + let views = elements.views(); + let buffers = data_buffers(&elements); + + let mut heads = HeadFilter::with_capacity(views.len()); + let mut short = HashTable::new(); + let mut long = HashTable::new(); + + for (idx, view) in views.iter().enumerate() { + heads.insert(view_head(view)); + + if view.is_inlined() { + let whole = view.as_u128(); + if let HashTableEntry::Vacant(vacant) = short.entry( + hasher.hash_one(whole), + |&other| other == whole, + |&other| hasher.hash_one(other), ) { - vacant.insert( - u32::try_from(idx).vortex_expect("a list holds fewer than 2^32 elements"), - ); + vacant.insert(whole); } + continue; + } + + let value = view.bytes(&buffers); + let bytes = |other: u32| views[other as usize].bytes(&buffers); + if let HashTableEntry::Vacant(vacant) = long.entry( + hasher.hash_one(value), + |&other| bytes(other) == value, + |&other| hasher.hash_one(bytes(other)), + ) { + vacant.insert( + u32::try_from(idx).vortex_expect("a list holds fewer than 2^32 elements"), + ); } } + Self { - elements, + heads, hasher, - table, + short, + elements, + long, } } + + /// Whether `view`, which points into `buffers`, is the view of an element. + #[inline] + fn contains_view( + &self, + view: &BinaryView, + buffers: &[&[u8]], + element_views: &[BinaryView], + element_buffers: &[&[u8]], + ) -> bool { + let head = view_head(view); + if !self.heads.may_contain(head) { + return false; + } + + if view.is_inlined() { + let whole = view.as_u128(); + return self + .short + .find(self.hasher.hash_one(whole), |&other| other == whole) + .is_some(); + } + + // The heads are equal, so the bytes can differ only after the 4-byte prefix. + let value = view.bytes(buffers); + self.long + .find(self.hasher.hash_one(value), |&idx| { + let element = &element_views[idx as usize]; + view_head(element) == head && element.bytes(element_buffers)[4..] == value[4..] + }) + .is_some() + } } impl Probe for BytesSet { @@ -250,20 +318,64 @@ impl Probe for BytesSet { let buffers = data_buffers(&array); let bits = collect_bits( array.views(), - |view: BinaryView| { - let value = view_bytes(&view, &buffers); - self.table - .find(self.hasher.hash_one(value), |&idx| { - view_bytes(&element_views[idx as usize], &element_buffers) == value - }) - .is_some() - }, + |view: BinaryView| self.contains_view(&view, &buffers, element_views, &element_buffers), ctx.allocator(), ); Ok((bits, array.validity()?)) } } +/// A filter of view heads with no false negatives: an inserted head always tests as present. +/// +/// It holds about [`HEAD_FILTER_BITS_PER_ELEMENT`] bits per element, and tests one bit, chosen by +/// a multiplicative hash of the head. +struct HeadFilter { + words: Box<[u64]>, + /// The shift that takes the top bits of the hash as the index of a bit. + shift: u32, +} + +impl HeadFilter { + fn with_capacity(elements: usize) -> Self { + let bits = (elements * HEAD_FILTER_BITS_PER_ELEMENT) + .next_power_of_two() + .max(64); + Self { + words: vec![0; bits / 64].into_boxed_slice(), + shift: u64::BITS - bits.trailing_zeros(), + } + } + + #[inline] + fn bit(&self, head: u64) -> usize { + // Fibonacci hashing: the top bits of the product depend on every bit of the head. + let hash = head.wrapping_mul(0x9E37_79B9_7F4A_7C15); + usize::try_from(hash >> self.shift).vortex_expect("a bit index fits a usize") + } + + fn insert(&mut self, head: u64) { + let bit = self.bit(head); + self.words[bit / 64] |= 1 << (bit % 64); + } + + #[inline] + fn may_contain(&self, head: u64) -> bool { + let bit = self.bit(head); + self.words[bit / 64] & (1 << (bit % 64)) != 0 + } +} + +/// The head of a view: the `u32` length of its value and the first 4 bytes of the value, +/// zero-padded for a value shorter than 4 bytes. +#[inline] +#[expect( + clippy::cast_possible_truncation, + reason = "the head is the low 8 bytes" +)] +fn view_head(view: &BinaryView) -> u64 { + view.as_u128() as u64 +} + /// Recursively canonical elements, indexed in sorted order with duplicates removed. struct RowSet { elements: ArrayRef, @@ -325,16 +437,6 @@ fn data_buffers(array: &VarBinViewArray) -> Vec<&[u8]> { .collect() } -/// The bytes a view points at, inlined in the view itself or out of line in one of `buffers`. -fn view_bytes<'a>(view: &'a BinaryView, buffers: &[&'a [u8]]) -> &'a [u8] { - if view.is_inlined() { - view.as_inlined().value() - } else { - let reference = view.as_view(); - &buffers[reference.buffer_index as usize][reference.as_range()] - } -} - /// The integer type with a float's bit pattern, or the type itself for an integer. fn bit_pattern_ptype(ptype: PType) -> PType { match ptype { diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs index 832580ccb17..c5ca8c3b174 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -21,6 +21,7 @@ use crate::arrays::FixedSizeListArray; use crate::arrays::ListArray; use crate::arrays::PrimitiveArray; use crate::arrays::StructArray; +use crate::arrays::VarBinViewArray; use crate::arrays::constant::list_scalar_elements; use crate::assert_arrays_eq; use crate::builders::builder_with_capacity_in; @@ -412,3 +413,46 @@ fn test_result_from_bits_rejects_mismatched_needles() -> VortexResult<()> { ); Ok(()) } + +#[test] +fn test_bytes_set_short_and_long_views() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let elements = VarBinViewArray::from_iter_nullable_str([ + Some("ab"), + Some(""), + Some("a long string with suffix A"), + Some("another long string"), + None, + Some("ab"), + ]) + .into_array(); + let set = PreparedSetData::try_new(elements, Nullability::NonNullable, &mut ctx)?; + + // "abc" shares the prefix of a short element, and "... suffix B" shares the head of a long + // one: length and first 4 bytes. Neither is an element. + let needles = VarBinViewArray::from_iter_nullable_str([ + Some("ab"), + Some("abc"), + Some(""), + Some("a long string with suffix B"), + Some("another long string"), + Some("a string never in the set"), + None, + ]) + .into_array(); + + assert_arrays_eq!( + set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, + BoolArray::from_iter([ + Some(true), + Some(false), + Some(true), + Some(false), + Some(true), + Some(false), + None, + ]), + &mut ctx + ); + Ok(()) +} From e98990ee96bd9863709c9ab1d369c59b1746aa12 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 09:52:35 -0400 Subject: [PATCH 8/9] Expose list membership with SQL null semantics in Java and Python Add Java/JNI list literals, listContains and IN/NOT IN builders. Expose the SQL null-semantics option and in_list in Python, including list literal coercion, public exports, typing, documentation and regression coverage. Signed-off-by: Robert Kruszewski --- docs/api/python/expr.rst | 3 + .../main/java/dev/vortex/api/Expression.java | 73 +++++++ .../java/dev/vortex/jni/NativeExpression.java | 6 + .../java/dev/vortex/api/ExpressionTest.java | 51 +++++ .../vortex/api/ListContainsFilterTest.java | 204 ++++++++++++++++++ vortex-jni/src/expression.rs | 149 +++++++++++-- vortex-python/python/vortex/_lib/expr.pyi | 7 +- vortex-python/python/vortex/expr.py | 2 + vortex-python/src/expr/mod.rs | 53 ++++- vortex-python/src/scalar/factory.rs | 49 +++-- vortex-python/test/test_expr.py | 42 ++++ vortex-python/test/test_scalar.py | 14 ++ 12 files changed, 615 insertions(+), 38 deletions(-) create mode 100644 java/vortex-jni/src/test/java/dev/vortex/api/ListContainsFilterTest.java diff --git a/docs/api/python/expr.rst b/docs/api/python/expr.rst index 9ee4a7d4450..17717387bba 100644 --- a/docs/api/python/expr.rst +++ b/docs/api/python/expr.rst @@ -52,6 +52,7 @@ Expressions are picklable, so a filter built in one process can be sent to anoth ~vortex.expr.pack ~vortex.expr.merge ~vortex.expr.list_contains + ~vortex.expr.in_list ~vortex.expr.list_length ~vortex.expr.list_sum ~vortex.expr.case_when @@ -153,6 +154,8 @@ Lists .. autofunction:: vortex.expr.list_contains +.. autofunction:: vortex.expr.in_list + .. autofunction:: vortex.expr.list_length .. autofunction:: vortex.expr.list_sum diff --git a/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java b/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java index b2e2a8be875..2608b2a41ba 100644 --- a/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java +++ b/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java @@ -154,6 +154,53 @@ public static Expression between( value.nativePointer(), lower.nativePointer(), upper.nativePointer(), lowerStrict, upperStrict)); } + /** + * Whether the list {@code list} evaluates to contains {@code needle}, with Vortex's null rules: a null element + * never matches anything, so a needle that matches no element is {@code false}. A null {@code needle} is null, + * except against an empty list, where the result is {@code false}. + * + *

{@code list} must evaluate to a Vortex list whose element type matches {@code needle}'s type, ignoring + * nullability. With a {@link #literalList(Expression...) list literal} on the left and a column on the right this + * is a set-membership test that a constant-set kernel answers in one pass over the column, however large the set. + * + * @see #listContains(Expression, Expression, boolean) + * @see #in(Expression, Expression) + */ + public static Expression listContains(Expression list, Expression needle) { + return listContains(list, needle, /* sqlNullSemantics= */ false); + } + + /** + * {@link #listContains(Expression, Expression)} with a choice of null rules. + * + * @param sqlNullSemantics {@code false} for Vortex's rules, where a null element never matches; {@code true} for + * SQL's three-valued {@code IN}, where a null element is an unknown value, so a needle that matches no element + * is {@code null} rather than {@code false} whenever the list holds a null. A null needle is null, except + * against an empty list with Vortex's default rules, where the result is {@code false}. + */ + public static Expression listContains(Expression list, Expression needle, boolean sqlNullSemantics) { + return new Expression( + NativeExpression.listContains(list.nativePointer(), needle.nativePointer(), sqlNullSemantics)); + } + + /** + * SQL {@code value IN (list)}: {@link #listContains(Expression, Expression, boolean)} with SQL null semantics and + * the operands in SQL order. A null {@code value} is null, and so is a non-match against a list that holds a null; + * both filter the row out of a scan. + */ + public static Expression in(Expression value, Expression list) { + return listContains(list, value, /* sqlNullSemantics= */ true); + } + + /** + * SQL {@code value NOT IN (list)}: the negation of {@link #in(Expression, Expression)}. Under SQL's rules a null + * {@code value}, or a list that holds a null, makes a non-match null rather than true, so such a row is never + * admitted. + */ + public static Expression notIn(Expression value, Expression list) { + return not(in(value, list)); + } + public static Expression literal(boolean value) { return new Expression(NativeExpression.literalBool(value, false)); } @@ -195,6 +242,32 @@ public static Expression literal(byte[] value) { return new Expression(NativeExpression.literalBinary(value)); } + /** + * Create a list literal out of literal element expressions, for example the right-hand side of an {@code IN} set. + * + *

Every element must itself be a literal and the elements must share a type, ignoring nullability; the list's + * element type is that shared type, made nullable if any element is. The list itself is non-null. + * + * @param elements at least one element; use {@link #literalEmptyList(DType)} for an empty list, which has no + * element to take a type from + */ + public static Expression literalList(Expression... elements) { + Preconditions.checkArgument( + elements.length > 0, + "literalList requires at least one element; use literalEmptyList for an empty list"); + return new Expression(NativeExpression.literalList(nativePointers(elements))); + } + + /** Create an empty list literal with the given nullable element type. It contains nothing, so nothing is in it. */ + public static Expression literalEmptyList(DType elementType) { + return new Expression(NativeExpression.literalEmptyList(elementType.tag(), false)); + } + + /** Create a null list literal with the given nullable element type. */ + public static Expression nullLiteralList(DType elementType) { + return new Expression(NativeExpression.literalEmptyList(elementType.tag(), true)); + } + /** * Create a decimal literal from its unscaled two's-complement big-endian byte representation (i.e. the value * returned by {@link BigInteger#toByteArray()}). diff --git a/java/vortex-jni/src/main/java/dev/vortex/jni/NativeExpression.java b/java/vortex-jni/src/main/java/dev/vortex/jni/NativeExpression.java index bcd82d4b313..8859357cddd 100644 --- a/java/vortex-jni/src/main/java/dev/vortex/jni/NativeExpression.java +++ b/java/vortex-jni/src/main/java/dev/vortex/jni/NativeExpression.java @@ -40,6 +40,8 @@ private NativeExpression() {} public static native long between( long valuePointer, long lowerPointer, long upperPointer, boolean lowerStrict, boolean upperStrict); + public static native long listContains(long listPointer, long needlePointer, boolean sqlNullSemantics); + public static native long literalBool(boolean value, boolean isNull); public static native long literalI8(byte value, boolean isNull); @@ -58,6 +60,10 @@ public static native long between( public static native long literalBinary(byte[] value); + public static native long literalList(long[] elementPointers); + + public static native long literalEmptyList(byte elementDTypeTag, boolean isNull); + public static native long literalDecimal(byte[] unscaledBigEndian, int precision, int scale, boolean isNull); public static native long literalDate(long value, byte timeUnitTag, boolean isNull); diff --git a/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTest.java b/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTest.java index e02024311d8..d0ab81e4de5 100644 --- a/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTest.java +++ b/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTest.java @@ -45,6 +45,57 @@ public void packComposes() { true)); } + @Test + public void literalListComposesWithListContains() { + Expression set = Expression.literalList(Expression.literal(1L), Expression.literal(2L)); + assertNotNull(Expression.listContains(set, Expression.column("id"))); + assertNotNull(Expression.in(Expression.column("id"), set)); + assertNotNull(Expression.notIn(Expression.column("id"), set)); + } + + @Test + public void literalListUnifiesElementNullability() { + // A null element makes the element type nullable rather than rejecting the set; the non-null elements are + // cast up to it. + assertNotNull(Expression.literalList( + Expression.literal(1L), Expression.nullLiteral(Expression.DType.I64), Expression.literal(3L))); + } + + @Test + public void literalListRequiresAtLeastOneElement() { + IllegalArgumentException exception = assertThrows(IllegalArgumentException.class, Expression::literalList); + assertTrue( + exception.getMessage().contains("literalEmptyList"), + () -> "unexpected message: " + exception.getMessage()); + } + + @Test + public void literalListRejectsMixedElementTypes() { + RuntimeException exception = assertThrows( + RuntimeException.class, + () -> Expression.literalList(Expression.literal(1L), Expression.literal("two"))); + assertTrue( + exception.getMessage().contains("must share a dtype"), + () -> "unexpected message: " + exception.getMessage()); + } + + @Test + public void literalListRejectsNonLiteralElements() { + RuntimeException exception = assertThrows( + RuntimeException.class, () -> Expression.literalList(Expression.literal(1L), Expression.column("id"))); + assertTrue( + exception.getMessage().contains("must themselves be literals"), + () -> "unexpected message: " + exception.getMessage()); + } + + @Test + public void emptyAndNullListLiteralsAcceptEveryNullLiteralDType() { + for (Expression.DType dtype : Expression.DType.values()) { + assertNotNull(Expression.literalEmptyList(dtype), () -> "native side rejected empty list of " + dtype); + assertNotNull(Expression.nullLiteralList(dtype), () -> "native side rejected null list of " + dtype); + } + } + @Test public void mergeComposes() { // Default duplicate handling (ERROR). diff --git a/java/vortex-jni/src/test/java/dev/vortex/api/ListContainsFilterTest.java b/java/vortex-jni/src/test/java/dev/vortex/api/ListContainsFilterTest.java new file mode 100644 index 00000000000..c04cc8800a1 --- /dev/null +++ b/java/vortex-jni/src/test/java/dev/vortex/api/ListContainsFilterTest.java @@ -0,0 +1,204 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +package dev.vortex.api; + +import static java.nio.charset.StandardCharsets.UTF_8; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import dev.vortex.arrow.ArrowAllocation; +import dev.vortex.jni.NativeLoader; +import java.io.IOException; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import org.apache.arrow.c.ArrowArray; +import org.apache.arrow.c.ArrowSchema; +import org.apache.arrow.c.Data; +import org.apache.arrow.memory.BufferAllocator; +import org.apache.arrow.vector.IntVector; +import org.apache.arrow.vector.VarCharVector; +import org.apache.arrow.vector.VectorSchemaRoot; +import org.apache.arrow.vector.ipc.ArrowReader; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.types.pojo.Schema; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +/** + * End-to-end coverage for {@link Expression#listContains(Expression, Expression)} as a scan filter. + * + *

Constructing a list-contains expression proves nothing on its own: the element type is only checked when the + * expression is bound to a schema, so a list literal built from the wrong element type builds fine and fails later, + * inside a scan. These tests therefore write a small file and read it back through the filter, asserting on the rows + * that survive. + * + *

The {@code maybe} column is nullable and null on every odd id, and the sets below hold a null element, to pin the + * two null rules apart: under SQL's, which {@link Expression#in} and {@link Expression#notIn} use, a null element makes + * every non-match null, so {@code NOT IN} keeps nothing; under Vortex's it is simply never a match. + * + *

The large-set case is the reason the binding exists. A caller without it has to expand {@code IN} into a chain of + * equality comparisons, which is why callers cap the set size and fall back to a bounding range; a single + * {@code list_contains} node carries the whole set, and the stats rewrite falsifies it for a zone only when every + * element misses that zone's bounds. + */ +public final class ListContainsFilterTest { + private static final int ROW_COUNT = 6; + + @TempDir + static Path tempDir; + + private static Session session; + private static String filePath; + + @BeforeAll + public static void loadLibrary() { + NativeLoader.loadJni(); + } + + @BeforeAll + static void writeFile() throws IOException { + session = Session.create(); + filePath = + tempDir.resolve("list_contains.vortex").toAbsolutePath().toUri().toString(); + + BufferAllocator allocator = ArrowAllocation.rootAllocator(); + Schema schema = new Schema(List.of( + Field.notNullable("id", new ArrowType.Int(32, true)), + Field.notNullable("name", new ArrowType.Utf8()), + Field.nullable("maybe", new ArrowType.Int(32, true)))); + + try (VortexWriter writer = VortexWriter.builder(session, filePath, schema, allocator) + .build(); + VectorSchemaRoot root = VectorSchemaRoot.create(schema, allocator)) { + IntVector id = (IntVector) root.getVector("id"); + VarCharVector name = (VarCharVector) root.getVector("name"); + IntVector maybe = (IntVector) root.getVector("maybe"); + id.allocateNew(ROW_COUNT); + name.allocateNew(ROW_COUNT); + maybe.allocateNew(ROW_COUNT); + for (int i = 0; i < ROW_COUNT; i++) { + id.setSafe(i, i + 1); + name.setSafe(i, ("row-" + (i + 1)).getBytes(UTF_8)); + // Every other row is null; the rest carry their id. + if (i % 2 == 0) { + maybe.setNull(i); + } else { + maybe.setSafe(i, i + 1); + } + } + root.setRowCount(ROW_COUNT); + + try (ArrowArray array = ArrowArray.allocateNew(allocator); + ArrowSchema arrowSchema = ArrowSchema.allocateNew(allocator)) { + Data.exportVectorSchemaRoot(allocator, root, null, array, arrowSchema); + writer.writeBatch(array.memoryAddress(), arrowSchema.memoryAddress()); + } + writer.finish(); + } + } + + @Test + public void inKeepsOnlyTheMatchingRows() { + Expression set = Expression.literalList(Expression.literal(2), Expression.literal(4), Expression.literal(99)); + assertEquals(List.of(2, 4), scanIds(Expression.in(Expression.column("id"), set))); + } + + @Test + public void notInKeepsTheComplement() { + Expression set = Expression.literalList(Expression.literal(2), Expression.literal(4), Expression.literal(99)); + assertEquals(List.of(1, 3, 5, 6), scanIds(Expression.notIn(Expression.column("id"), set))); + } + + @Test + public void listContainsTakesTheListFirst() { + // in(value, list) is listContains(list, value) with the operands the other way round; both spellings have to + // agree, and passing the column as the list would be a type error rather than a silently different filter. + Expression set = Expression.literalList(Expression.literal(3)); + assertEquals(List.of(3), scanIds(Expression.listContains(set, Expression.column("id")))); + } + + @Test + public void inWithANullElementKeepsOnlyMatches() { + // maybe = [null, 2, null, 4, null, 6]; the set holds 2 and a null. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(2), scanIds(Expression.in(Expression.column("maybe"), set))); + } + + @Test + public void notInWithANullElementKeepsNothing() { + // SQL: x NOT IN (2, NULL) is false for 2 and null for everything else, so no row survives. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(), scanIds(Expression.notIn(Expression.column("maybe"), set))); + // Without the null element NOT IN keeps the non-null non-matches. + assertEquals( + List.of(4, 6), + scanIds(Expression.notIn(Expression.column("maybe"), Expression.literalList(Expression.literal(2))))); + } + + @Test + public void vortexNullRulesTreatANullElementAsNoMatch() { + // listContains without SQL semantics: the null element never matches, so a non-match stays false and its + // negation keeps the non-null non-matches. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(2), scanIds(Expression.listContains(set, Expression.column("maybe")))); + assertEquals(List.of(4, 6), scanIds(Expression.not(Expression.listContains(set, Expression.column("maybe"))))); + } + + @Test + public void aStringSetFiltersOnUtf8() { + Expression set = Expression.literalList(Expression.literal("row-1"), Expression.literal("row-6")); + assertEquals(List.of(1, 6), scanIds(Expression.in(Expression.column("name"), set))); + } + + @Test + public void aSetLargerThanAnyOrChainCapStillPushesDown() { + // The whole point of the binding: 1000 literals stay one expression node. Only 5 is present in the file. + Expression[] elements = new Expression[1000]; + for (int i = 0; i < elements.length; i++) { + elements[i] = Expression.literal(i == 0 ? 5 : ROW_COUNT + i); + } + assertEquals(List.of(5), scanIds(Expression.in(Expression.column("id"), Expression.literalList(elements)))); + } + + @Test + public void anEmptySetMatchesNothing() { + Expression empty = Expression.literalEmptyList(Expression.DType.I32); + assertTrue(scanIds(Expression.in(Expression.column("id"), empty)).isEmpty()); + } + + @Test + public void aNullSetMatchesNothing() { + // A null list yields null rather than false, which filters the row out just the same. + Expression nullList = Expression.nullLiteralList(Expression.DType.I32); + assertTrue(scanIds(Expression.in(Expression.column("id"), nullList)).isEmpty()); + } + + /** Reads the {@code id} column of every row that survives {@code filter}, in file order. */ + private static List scanIds(Expression filter) { + BufferAllocator allocator = ArrowAllocation.rootAllocator(); + DataSource dataSource = DataSource.open(session, filePath); + Scan scan = dataSource.scan( + ScanOptions.builder().filter(filter).ordered(true).build()); + + List ids = new ArrayList<>(); + while (scan.hasNext()) { + Partition partition = scan.next(); + try (ArrowReader reader = partition.scanArrow(allocator)) { + while (reader.loadNextBatch()) { + VectorSchemaRoot root = reader.getVectorSchemaRoot(); + IntVector id = (IntVector) root.getVector("id"); + for (int i = 0; i < root.getRowCount(); i++) { + ids.add(id.get(i)); + } + } + } catch (IOException e) { + throw new AssertionError("failed reading partition", e); + } + } + return ids; + } +} diff --git a/vortex-jni/src/expression.rs b/vortex-jni/src/expression.rs index 77d709ddd38..154f7056620 100644 --- a/vortex-jni/src/expression.rs +++ b/vortex-jni/src/expression.rs @@ -38,6 +38,7 @@ use vortex::expr::between; use vortex::expr::get_item; use vortex::expr::is_not_null; use vortex::expr::is_null; +use vortex::expr::list_contains_opts; use vortex::expr::lit; use vortex::expr::merge_opts; use vortex::expr::not; @@ -60,6 +61,8 @@ use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; use vortex::scalar_fn::fns::like::Like; use vortex::scalar_fn::fns::like::LikeOptions; +use vortex::scalar_fn::fns::list_contains::ListContainsOptions; +use vortex::scalar_fn::fns::literal::Literal; use vortex::scalar_fn::fns::merge::DuplicateHandling; use vortex::scalar_fn::fns::operators::Operator; @@ -362,6 +365,33 @@ fn strict_from_bool(value: jboolean) -> StrictComparison { } } +/// Build `list_contains(list, needle)`: whether the list-typed `list` expression contains +/// `needle`. +/// +/// `list` must evaluate to a Vortex `List`; the list's element dtype must match `needle`'s dtype +/// ignoring nullability. With a list literal on the left and a column on the right this is a +/// set-membership (`IN`) test that a constant-set kernel answers in one pass over the column. +/// +/// `sql_null_semantics` selects how a null list element behaves: off, it never matches and a +/// non-matching needle is `false`; on, it is SQL's unknown value and a non-matching needle is +/// `null`, so `NOT IN` never admits it. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_listContains( + _env: EnvUnowned, + _class: JClass, + list: jlong, + needle: jlong, + sql_null_semantics: jboolean, +) -> jlong { + let list = unsafe { expr_ref(list) }.clone(); + let needle = unsafe { expr_ref(needle) }.clone(); + into_raw(list_contains_opts( + list, + needle, + ListContainsOptions { sql_null_semantics }, + )) +} + #[unsafe(no_mangle)] pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalBool( _env: EnvUnowned, @@ -437,6 +467,93 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalBinary( }) } +/// Build a non-empty list literal out of the literal expressions in `elements`. +/// +/// Every element must be a literal (`vortex.literal`) whose dtype matches the first element's +/// ignoring nullability; the list's element dtype is that shared dtype, made nullable if any +/// element is nullable, and each element is cast to it. The list itself is non-nullable — use +/// [`Java_dev_vortex_jni_NativeExpression_literalEmptyList`] for a null or empty list, which +/// cannot infer an element dtype from its (absent) elements. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalList( + mut env: EnvUnowned, + _class: JClass, + elements: JLongArray, +) -> jlong { + try_or_throw(&mut env, |env| { + let ptrs = unsafe { elements.get_elements(env, ReleaseMode::NoCopyBack) }?; + let scalars = ptrs + .iter() + .map(|ptr| literal_scalar(unsafe { expr_ref(*ptr) })) + .collect::, _>>()?; + Ok(into_raw(lit(list_scalar(&scalars)?))) + }) +} + +/// The scalar behind a literal expression, or an error if the expression is not a literal. +fn literal_scalar(expr: &Expression) -> Result { + expr.as_opt::().cloned().ok_or_else(|| { + vortex_err!("list literal elements must themselves be literals, got {expr}").into() + }) +} + +/// Collect literal scalars into a single non-nullable list scalar. +fn list_scalar(elements: &[Scalar]) -> Result { + let Some(first) = elements.first() else { + throw_runtime!("list literal requires at least one element; use an empty list literal"); + }; + + let mut nullability = Nullability::NonNullable; + for element in elements { + if !element.dtype().eq_ignore_nullability(first.dtype()) { + throw_runtime!( + "list literal elements must share a dtype, got {} and {}", + first.dtype(), + element.dtype() + ); + } + nullability |= element.dtype().nullability(); + } + + let element_dtype = first.dtype().with_nullability(nullability); + let children = elements + .iter() + .map(|element| element.cast(&element_dtype)) + .collect::, _>>()?; + Ok(Scalar::list( + Arc::new(element_dtype), + children, + Nullability::NonNullable, + )) +} + +/// Build an empty (or null) list literal whose element dtype is selected by `element_dtype_tag`. +/// +/// The tag table is the one [`Java_dev_vortex_jni_NativeExpression_literalNull`] reads; see +/// `dev.vortex.api.Expression.DType` on the Java side for the source of truth. Elements are +/// nullable so that the literal accepts a nullable column as its needle. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalEmptyList( + mut env: EnvUnowned, + _class: JClass, + element_dtype_tag: jbyte, + is_null_flag: jboolean, +) -> jlong { + try_or_throw(&mut env, |_| { + let element_dtype = Arc::new(parse_null_dtype(element_dtype_tag)?); + if is_null_flag { + return Ok(into_raw(lit(Scalar::null(DType::List( + element_dtype, + Nullability::Nullable, + ))))); + } + Ok(into_raw(lit(Scalar::list_empty( + element_dtype, + Nullability::NonNullable, + )))) + }) +} + /// Build a decimal literal from a two's-complement big-endian byte representation of the /// unscaled value (the format produced by Java's `BigInteger.toByteArray()`). #[unsafe(no_mangle)] @@ -666,10 +783,26 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalUuid( }) } -/// Build a typed null literal whose nullable dtype is selected by `dtype_tag`. +/// Parse a nullable primitive [`DType`] from the wire-encoded byte tag. /// /// Tag values intentionally do not overlap with [`parse_time_unit`]. /// See `dev.vortex.api.Expression.DType` on the Java side for the source of truth. +fn parse_null_dtype(tag: jbyte) -> Result { + Ok(match tag { + 0 => DType::Bool(Nullability::Nullable), + 1 => DType::Primitive(PType::I8, Nullability::Nullable), + 2 => DType::Primitive(PType::I16, Nullability::Nullable), + 3 => DType::Primitive(PType::I32, Nullability::Nullable), + 4 => DType::Primitive(PType::I64, Nullability::Nullable), + 5 => DType::Primitive(PType::F32, Nullability::Nullable), + 6 => DType::Primitive(PType::F64, Nullability::Nullable), + 7 => DType::Utf8(Nullability::Nullable), + 8 => DType::Binary(Nullability::Nullable), + other => throw_runtime!("unknown null dtype tag: {other}"), + }) +} + +/// Build a typed null literal whose nullable dtype is selected by `dtype_tag`. #[unsafe(no_mangle)] pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalNull( mut env: EnvUnowned, @@ -677,18 +810,6 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalNull( dtype_tag: jbyte, ) -> jlong { try_or_throw(&mut env, |_| { - let dtype = match dtype_tag { - 0 => DType::Bool(Nullability::Nullable), - 1 => DType::Primitive(PType::I8, Nullability::Nullable), - 2 => DType::Primitive(PType::I16, Nullability::Nullable), - 3 => DType::Primitive(PType::I32, Nullability::Nullable), - 4 => DType::Primitive(PType::I64, Nullability::Nullable), - 5 => DType::Primitive(PType::F32, Nullability::Nullable), - 6 => DType::Primitive(PType::F64, Nullability::Nullable), - 7 => DType::Utf8(Nullability::Nullable), - 8 => DType::Binary(Nullability::Nullable), - other => throw_runtime!("unknown null dtype tag: {other}"), - }; - Ok(into_raw(lit(Scalar::null(dtype)))) + Ok(into_raw(lit(Scalar::null(parse_null_dtype(dtype_tag)?)))) }) } diff --git a/vortex-python/python/vortex/_lib/expr.pyi b/vortex-python/python/vortex/_lib/expr.pyi index 6cd2b58bfd5..9d5e7be962a 100644 --- a/vortex-python/python/vortex/_lib/expr.pyi +++ b/vortex-python/python/vortex/_lib/expr.pyi @@ -8,9 +8,9 @@ from typing import Literal, TypeAlias, final from typing_extensions import override from .dtype import DType -from .scalar import ScalarPyType +from .scalar import Scalar, ScalarPyType -IntoExpr: TypeAlias = Expr | bool | int | float | str | bytes | date | datetime | None +IntoExpr: TypeAlias = Expr | Scalar | ScalarPyType | date | datetime """A value accepted anywhere an expression is expected. Non-``Expr`` values become literals.""" VariantPath: TypeAlias = str | int | Sequence[str | int] @@ -103,7 +103,8 @@ def merge( ) -> Expr: ... # Lists -def list_contains(child: IntoExpr, value: IntoExpr) -> Expr: ... +def list_contains(child: IntoExpr, value: IntoExpr, *, sql_null_semantics: bool = False) -> Expr: ... +def in_list(value: IntoExpr, list: IntoExpr) -> Expr: ... def list_length(child: IntoExpr) -> Expr: ... def list_sum(child: IntoExpr, *, skip_nans: bool = True) -> Expr: ... diff --git a/vortex-python/python/vortex/expr.py b/vortex-python/python/vortex/expr.py index d02ed5ebab6..afd78cb4d85 100644 --- a/vortex-python/python/vortex/expr.py +++ b/vortex-python/python/vortex/expr.py @@ -21,6 +21,7 @@ gt, gt_eq, ilike, + in_list, is_not_null, is_null, like, @@ -67,6 +68,7 @@ "gt", "gt_eq", "ilike", + "in_list", "is_not_null", "is_null", "like", diff --git a/vortex-python/src/expr/mod.rs b/vortex-python/src/expr/mod.rs index 8d76d20b051..03f76042a93 100644 --- a/vortex-python/src/expr/mod.rs +++ b/vortex-python/src/expr/mod.rs @@ -23,6 +23,7 @@ use vortex::scalar_fn::ScalarFnVTableExt; use vortex::scalar_fn::fns::between::BetweenOptions; use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; +use vortex::scalar_fn::fns::list_contains::ListContainsOptions; use vortex::scalar_fn::fns::merge::DuplicateHandling; use vortex::scalar_fn::fns::operators::Operator; use vortex::scalar_fn::fns::variant_get::VariantPath; @@ -85,6 +86,7 @@ pub(crate) fn init(py: Python, parent: &Bound) -> PyResult<()> { // Lists m.add_function(wrap_pyfunction!(list_contains, &m)?)?; + m.add_function(wrap_pyfunction!(in_list, &m)?)?; m.add_function(wrap_pyfunction!(list_length, &m)?)?; m.add_function(wrap_pyfunction!(list_sum, &m)?)?; @@ -502,7 +504,8 @@ pub fn literal<'py>( dtype: &Bound<'py, PyDType>, value: &Bound<'py, PyAny>, ) -> PyResult> { - scalar(dtype.borrow().inner().clone(), value) + let scalar = scalar_helper(value, Some(dtype.borrow().inner())).map_err(PyErr::from)?; + Bound::new(value.py(), PyExpr { inner: lit(scalar) }) } /// Create an expression that refers to the identity scope. @@ -1097,20 +1100,64 @@ pub fn merge(exprs: &Bound<'_, PyAny>, duplicate_handling: &str) -> PyResult PyExpr { + PyExpr { + inner: expr::list_contains_opts( + child.into_inner(), + value.into_inner(), + ListContainsOptions { sql_null_semantics }, + ), + } +} + +/// SQL ``value IN (list)`` with SQL null semantics. +/// +/// A null value produces null. A non-match also produces null if the list contains null. +/// Use ``~in_list(value, list)`` for SQL ``NOT IN``. +/// +/// Parameters +/// ---------- +/// value : :class:`Any` +/// The value to search for. +/// list : :class:`Any` +/// A list expression or a Python list. Use :func:`.literal` with an explicit list dtype for +/// empty or null lists and for element types other than the inferred Python scalar types. /// /// Returns /// ------- /// :class:`vortex.Expr` +/// +/// Examples +/// -------- +/// +/// ```python +/// >>> import vortex.expr as ve +/// >>> ve.in_list(ve.column("age"), [25, 30]) +/// +/// ``` #[pyfunction] -pub fn list_contains(child: PyIntoExpr, value: PyIntoExpr) -> PyExpr { +#[pyo3(signature = (value, list))] +pub fn in_list(value: PyIntoExpr, list: PyIntoExpr) -> PyExpr { PyExpr { - inner: expr::list_contains(child.into_inner(), value.into_inner()), + inner: expr::in_list(value.into_inner(), list.into_inner()), } } diff --git a/vortex-python/src/scalar/factory.rs b/vortex-python/src/scalar/factory.rs index 13e46b048c9..968e45e55c9 100644 --- a/vortex-python/src/scalar/factory.rs +++ b/vortex-python/src/scalar/factory.rs @@ -22,6 +22,7 @@ use vortex::scalar::DecimalValue; use vortex::scalar::Scalar; use crate::dtype::PyDType; +use crate::error::PyVortexError; use crate::error::PyVortexResult; use crate::scalar::PyScalar; use crate::scalar::bool; @@ -180,32 +181,44 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes if let Some(DType::List(element_dtype, ..)) = dtype { let elements = list .iter() - .map(|e| scalar_helper_inner(&e, Some(element_dtype))) - .try_collect()?; - Scalar::list( + .map(|e| scalar_helper(&e, Some(element_dtype))) + .collect::>>()?; + return Ok(Scalar::list( Arc::clone(element_dtype), elements, Nullability::NonNullable, - ); + )); } else { - // If no dtype was provided, we need to infer the element dtype from the list contents. - // We do this in a greedy way taking the first element dtype we find. - let mut elements = Vec::with_capacity(list.len()); - let mut element_dtype = None; - - for element in list.iter() { - let scalar = scalar_helper_inner(&element, element_dtype.as_ref())?; - if element_dtype.is_none() { - element_dtype = Some(scalar.dtype().clone()); + let elements = list + .iter() + .map(|element| scalar_helper_inner(&element, None)) + .collect::>>()?; + let element_dtype = elements + .iter() + .find(|element| !matches!(element.dtype(), DType::Null)) + .map(|element| element.dtype().clone()) + .unwrap_or(DType::Null); + let mut nullability = element_dtype.nullability(); + for element in &elements { + if !matches!(element.dtype(), DType::Null) + && !element.dtype().eq_ignore_nullability(&element_dtype) + { + return Err(PyValueError::new_err(format!( + "list elements must share a dtype, got {} and {}", + element_dtype, + element.dtype() + ))); } - elements.push(scalar); + nullability |= element.dtype().nullability(); } + let element_dtype = element_dtype.with_nullability(nullability); + let elements = elements + .iter() + .map(|element| element.cast(&element_dtype).map_err(PyVortexError::from)) + .collect::>>()?; return Ok(Scalar::list( - element_dtype - .map(Arc::new) - // Empty list defaults to Null dtype - .unwrap_or_else(|| Arc::new(DType::Null)), + Arc::new(element_dtype), elements, Nullability::NonNullable, )); diff --git a/vortex-python/test/test_expr.py b/vortex-python/test/test_expr.py index ed5cf3bee76..959de013b55 100644 --- a/vortex-python/test/test_expr.py +++ b/vortex-python/test/test_expr.py @@ -84,6 +84,8 @@ def column_values(vxf: vx.VortexFile, projection: Expr) -> list[object]: "merge": lambda: ve.merge([ve.select(["name"]), ve.select(["age"])]), "merge_rightmost": lambda: ve.merge([ve.select(["name"]), ve.select(["name"])], duplicate_handling="rightmost"), "list_contains": lambda: ve.list_contains(ve.column("scores"), 5), + "list_contains_sql": lambda: ve.list_contains(ve.column("scores"), 5, sql_null_semantics=True), + "in_list": lambda: ve.in_list(ve.column("age"), [25, 30]), "list_length": lambda: ve.list_length(ve.column("scores")), "list_sum": lambda: ve.list_sum(ve.column("scores")), "list_sum_nans": lambda: ve.list_sum(ve.column("scores"), skip_nans=False), @@ -209,6 +211,46 @@ def test_arithmetic_and_list_functions(people: vx.VortexFile) -> None: assert column_values(people, ve.get_item("city", ve.column("nested"))) == ["Paris", "Berlin", "Paris", "Lima"] +@pytest.mark.parametrize( + "values,default,sql", + [ + ([30, 57], [True, False, None, True], [True, False, None, True]), + ([30, None], [True, False, None, False], [True, None, None, None]), + ([None, 30], [True, False, None, False], [True, None, None, None]), + ([None], [False, False, None, False], [None, None, None, None]), + ([], [False, False, False, False], [False, False, None, False]), + (None, [None, None, None, None], [None, None, None, None]), + ], +) +def test_list_membership_null_semantics( + people: vx.VortexFile, + values: list[int | None] | None, + default: list[bool | None], + sql: list[bool | None], +): + members = ve.literal(vx.list_(vx.int_(64, nullable=True), nullable=True), values) + age = ve.column("age") + assert column_values(people, ve.list_contains(members, age)) == default + assert column_values(people, ve.list_contains(members, age, sql_null_semantics=True)) == sql + expr = ve.in_list(age, members) + assert column_values(people, expr) == sql + assert column_values(people, ve.deserialize(expr.serialize())) == sql + assert column_values(people, ~expr) == [None if value is None else not value for value in sql] + + +@pytest.mark.parametrize("members", [[30, None], [None, 30]]) +def test_in_list_python_list_filter(people: vx.VortexFile, members: list[int | None]): + expr = ve.in_list(ve.column("age"), members) + assert names(people, expr) == ["Alice"] + assert names(people, ~expr) == [] + assert names(people, ~ve.in_list(ve.column("age"), [30])) == ["Bob", "Charlie"] + + +def test_in_list_large_set_and_strings(people: vx.VortexFile): + assert names(people, ve.in_list(ve.column("age"), list(range(1000)))) == ["Alice", "Bob", "Charlie"] + assert names(people, ve.in_list(ve.column("name"), ["Alice", "Charlie"])) == ["Alice", "Charlie"] + + def test_case_when_semantics(people: vx.VortexFile) -> None: expr = ve.case_when([(ve.gt(ve.column("age"), 40), "senior"), (ve.gt(ve.column("age"), 26), "mid")], "junior") assert column_values(people, expr) == ["mid", "junior", "junior", "senior"] diff --git a/vortex-python/test/test_scalar.py b/vortex-python/test/test_scalar.py index 24c5889c9a4..9af210dc34b 100644 --- a/vortex-python/test/test_scalar.py +++ b/vortex-python/test/test_scalar.py @@ -39,6 +39,20 @@ def test_f16() -> None: assert scalar.as_py() == 1.0 +@pytest.mark.parametrize("values", [[1, None], [None, 1], [None, None], []]) +def test_list_scalar_nullability(values: list[int | None]): + assert vx.scalar(values).as_py() == values + dtype = vx.list_(vx.int_(32, nullable=True)) + scalar = vx.scalar(values, dtype=dtype) + assert scalar.dtype == dtype + assert scalar.as_py() == values + + +def test_list_scalar_rejects_mixed_types(): + with pytest.raises(ValueError, match="must share a dtype"): + _ = vx.scalar([1, "two"]) + + @pytest.mark.parametrize( "precision,scale,stored,expected", [ From 2a516129912c8591ce808c8fa0c0e5fdd90d37d4 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 11:52:15 -0400 Subject: [PATCH 9/9] Fix Python expression binding compilation and typing Signed-off-by: Robert Kruszewski --- vortex-python/python/vortex/_lib/expr.pyi | 6 +++--- vortex-python/src/expr/mod.rs | 11 ----------- 2 files changed, 3 insertions(+), 14 deletions(-) diff --git a/vortex-python/python/vortex/_lib/expr.pyi b/vortex-python/python/vortex/_lib/expr.pyi index 9d5e7be962a..7544b6f830d 100644 --- a/vortex-python/python/vortex/_lib/expr.pyi +++ b/vortex-python/python/vortex/_lib/expr.pyi @@ -3,14 +3,14 @@ from collections.abc import Iterable, Mapping, Sequence from datetime import date, datetime -from typing import Literal, TypeAlias, final +from typing import Any, Literal, TypeAlias, final from typing_extensions import override from .dtype import DType from .scalar import Scalar, ScalarPyType -IntoExpr: TypeAlias = Expr | Scalar | ScalarPyType | date | datetime +IntoExpr: TypeAlias = Expr | Scalar | ScalarPyType | list[Any] | dict[str, Any] | date | datetime """A value accepted anywhere an expression is expected. Non-``Expr`` values become literals.""" VariantPath: TypeAlias = str | int | Sequence[str | int] @@ -46,7 +46,7 @@ class Expr: # Leaves and scope def root() -> Expr: ... def column(name: str) -> Expr: ... -def literal(dtype: DType, value: ScalarPyType) -> Expr: ... +def literal(dtype: DType, value: object) -> Expr: ... def get_item(field: str, child: IntoExpr | None = None) -> Expr: ... # Boolean logic diff --git a/vortex-python/src/expr/mod.rs b/vortex-python/src/expr/mod.rs index 03f76042a93..806aa5ddcd5 100644 --- a/vortex-python/src/expr/mod.rs +++ b/vortex-python/src/expr/mod.rs @@ -10,7 +10,6 @@ use pyo3::intern; use pyo3::prelude::*; use pyo3::types::*; use vortex::aggregate_fn::NumericalAggregateOpts; -use vortex::dtype::DType; use vortex::dtype::FieldName; use vortex::dtype::FieldNames; use vortex::dtype::Nullability; @@ -596,16 +595,6 @@ pub fn get_item(field: String, child: Option) -> PyExpr { } } -pub fn scalar<'py>(dtype: DType, value: &Bound<'py, PyAny>) -> PyResult> { - let py = value.py(); - Bound::new( - py, - PyExpr { - inner: lit(scalar_helper(value, Some(&dtype))?), - }, - ) -} - /// Negate a Boolean expression. /// /// Parameters