From f3e5ee01e36b9021c0627374b6dc8f44f7afdd12 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 10:09:14 -0400 Subject: [PATCH 01/10] 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 983a89ba714823759b54b668d7af5c9bcb831a40 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 11:01:31 -0400 Subject: [PATCH 02/10] 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 c8848368c51ba06ea6695d159b4712d861b0a3d2 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Fri, 25 Sep 2026 12:27:52 -0400 Subject: [PATCH 03/10] 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 09547e7afd5b7e6779c9534e433212d7464fcdbf Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 18:28:12 -0400 Subject: [PATCH 04/10] 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 cdee167ea71..2eac08dfa9d 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 e4c0339c6c61f94925e14f46937f2b242e419716 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 22:11:20 -0400 Subject: [PATCH 05/10] 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 a351995c0852a8ad83e7616ac816eb9f6ba50d3a Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Mon, 28 Sep 2026 22:51:00 -0400 Subject: [PATCH 06/10] 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 77f114d0770d3374d4c072eafd9f3817de71dff0 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Tue, 29 Sep 2026 12:12:56 -0700 Subject: [PATCH 07/10] 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 843f9c8f090a0b6f24c29805317dfabe528f35d2 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Tue, 29 Sep 2026 15:24:33 -0700 Subject: [PATCH 08/10] more Signed-off-by: Robert Kruszewski --- .../src/arrays/constant/vtable/canonical.rs | 201 +++++++++++++- .../src/scalar_fn/fns/list_contains/mod.rs | 159 +++++++++-- .../fns/list_contains/prepared/array.rs | 259 ++++++++++-------- .../fns/list_contains/prepared/literal.rs | 104 +++++++ .../fns/list_contains/prepared/mod.rs | 8 +- .../fns/list_contains/prepared/tests.rs | 62 +++-- vortex-array/src/scalar_fn/session.rs | 2 + 7 files changed, 611 insertions(+), 184 deletions(-) create mode 100644 vortex-array/src/scalar_fn/fns/list_contains/prepared/literal.rs diff --git a/vortex-array/src/arrays/constant/vtable/canonical.rs b/vortex-array/src/arrays/constant/vtable/canonical.rs index 7c5ff18106f..787aa028303 100644 --- a/vortex-array/src/arrays/constant/vtable/canonical.rs +++ b/vortex-array/src/arrays/constant/vtable/canonical.rs @@ -7,11 +7,13 @@ 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; use vortex_error::VortexExpect; use vortex_error::VortexResult; +use vortex_error::vortex_panic; use crate::ArrayRef; use crate::Canonical; @@ -34,6 +36,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; @@ -45,6 +49,7 @@ 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. @@ -311,23 +316,116 @@ pub(crate) fn list_scalar_elements(list: &ListScalar, allocator: &BufferAllocato let Some(elements) = list.element_values() else { return Canonical::empty(element_dtype).into_array(); }; + flat_elements(element_dtype, elements, allocator) +} - let mut builder = builder_with_capacity_in(element_dtype, elements.len(), allocator); - 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"); +/// The array of `elements`, the values of a list of `element_dtype`. +/// +/// Flat elements are written straight from their values, since a scalar per element costs more +/// than a probe over the result does. Nested elements go through the canonical builder. +fn flat_elements( + element_dtype: &DType, + elements: &[Option], + allocator: &BufferAllocatorRef, +) -> ArrayRef { + match element_dtype { + DType::Null => NullArray::new(elements.len()).into_array(), + DType::Bool(nullability) => { + let bits = BitBuffer::from_iter(elements.iter().map(|element| match element { + Some(ScalarValue::Bool(value)) => *value, + Some(_) => vortex_panic!("bool list holds a non-bool element"), + None => false, + })); + BoolArray::new(bits, element_validity(elements, *nullability)).into_array() + } + DType::Primitive(ptype, nullability) => match_each_native_ptype!(ptype, |T| { + let mut buffer = BufferMut::::with_capacity_in(elements.len(), allocator.clone()); + buffer.extend(elements.iter().map(|element| { + match element { + Some(ScalarValue::Primitive(value)) => value + .cast::() + .vortex_expect("list element of the list's element ptype"), + Some(_) => vortex_panic!("primitive list holds a non-primitive element"), + None => T::default(), + } + })); + PrimitiveArray::new(buffer.freeze(), element_validity(elements, *nullability)) + .into_array() + }), + DType::Decimal(decimal, nullability) => { + let values_type = DecimalType::smallest_decimal_value_type(decimal); + match_each_decimal_value_type!(values_type, |T| { + let mut buffer = + BufferMut::::with_capacity_in(elements.len(), allocator.clone()); + buffer.extend(elements.iter().map(|element| { + match element { + Some(ScalarValue::Decimal(value)) => value + .cast::() + .vortex_expect("list element fits the list's decimal type"), + Some(_) => vortex_panic!("decimal list holds a non-decimal element"), + None => T::default(), + } + })); + DecimalArray::new( + buffer.freeze(), + *decimal, + element_validity(elements, *nullability), + ) + .into_array() + }) + } + DType::Utf8(_) | DType::Binary(_) => { + let mut builder = VarBinViewBuilder::with_capacity_in( + element_dtype.clone(), + elements.len(), + allocator.clone(), + ); + for element in elements { + match element { + None => builder.append_null(), + Some(ScalarValue::Utf8(value)) => builder.append_value(value.as_bytes()), + Some(ScalarValue::Binary(value)) => builder.append_value(value.as_slice()), + Some(_) => vortex_panic!("byte list holds a non-byte element"), + } } - None => { - builder.append_null(); + builder.finish() + } + // An extension value is its storage value. + DType::Extension(ext_dtype) => ExtensionArray::new( + ext_dtype.clone(), + flat_elements(ext_dtype.storage_dtype(), elements, allocator), + ) + .into_array(), + _ => { + let mut builder = builder_with_capacity_in(element_dtype, elements.len(), allocator); + 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() + } + } +} + +/// The validity of a list's elements, a null element being a `None` value. +fn element_validity(elements: &[Option], nullability: Nullability) -> Validity { + match nullability { + Nullability::NonNullable => Validity::NonNullable, + Nullability::Nullable if elements.iter().all(Option::is_some) => Validity::AllValid, + Nullability::Nullable => { + Validity::from(BitBuffer::from_iter(elements.iter().map(Option::is_some))) } } - builder.finish() } /// Creates a [`FixedSizeListArray`] whose every row holds the same list. @@ -413,14 +511,20 @@ mod tests { use rstest::rstest; use vortex_error::VortexExpect; use vortex_error::VortexResult; + use vortex_error::vortex_panic; use vortex_session::VortexSession; + use super::list_scalar_elements; + use crate::ArrayRef; use crate::Canonical; use crate::IntoArray; use crate::VortexSessionExecute; + use crate::arrays::BoolArray; use crate::arrays::Chunked; use crate::arrays::Constant; use crate::arrays::ConstantArray; + use crate::arrays::DecimalArray; + use crate::arrays::ExtensionArray; use crate::arrays::FixedSizeListArray; use crate::arrays::ListViewArray; use crate::arrays::NullArray; @@ -436,11 +540,14 @@ mod tests { use crate::arrays::struct_::StructArrayExt; use crate::assert_arrays_eq; use crate::dtype::DType; + use crate::dtype::DecimalDType; use crate::dtype::Nullability; use crate::dtype::PType; use crate::dtype::half::f16; use crate::expr::stats::Stat; use crate::expr::stats::StatsProvider; + use crate::extension::datetime::Date; + use crate::extension::datetime::TimeUnit; use crate::scalar::Scalar; use crate::validity::Validity; @@ -1029,4 +1136,74 @@ mod tests { &mut ctx ); } + + /// Every flat element dtype is written straight from its values, null elements included. + #[rstest] + #[case::bool( + Scalar::list( + Arc::new(DType::Bool(Nullability::Nullable)), + vec![ + Scalar::bool(true, Nullability::Nullable), + Scalar::null(DType::Bool(Nullability::Nullable)), + Scalar::bool(false, Nullability::Nullable), + ], + Nullability::NonNullable, + ), + BoolArray::from_iter([Some(true), None, Some(false)]).into_array() + )] + #[case::decimal( + Scalar::list( + Arc::new(DType::Decimal(DecimalDType::new(5, 2), Nullability::Nullable)), + vec![ + Scalar::decimal(123i32.into(), DecimalDType::new(5, 2), Nullability::Nullable), + Scalar::null(DType::Decimal(DecimalDType::new(5, 2), Nullability::Nullable)), + Scalar::decimal((-45i32).into(), DecimalDType::new(5, 2), Nullability::Nullable), + ], + Nullability::NonNullable, + ), + DecimalArray::from_option_iter::([Some(123), None, Some(-45)], DecimalDType::new(5, 2)) + .into_array() + )] + #[case::null( + Scalar::list( + Arc::new(DType::Null), + vec![Scalar::null(DType::Null), Scalar::null(DType::Null)], + Nullability::NonNullable, + ), + NullArray::new(2).into_array() + )] + #[case::extension(date_list(), date_elements())] + fn test_list_scalar_elements_of_flat_dtypes( + #[case] list: Scalar, + #[case] expected: ArrayRef, + ) -> VortexResult<()> { + let mut ctx = SESSION.create_execution_ctx(); + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + assert_arrays_eq!(elements, expected, &mut ctx); + Ok(()) + } + + /// `[date(42), null]` as a list scalar. + fn date_list() -> Scalar { + let date = Scalar::extension::(TimeUnit::Days, Scalar::from(Some(42i32))); + let dtype = date.dtype().clone(); + Scalar::list( + Arc::new(dtype.clone()), + vec![date, Scalar::null(dtype)], + Nullability::NonNullable, + ) + } + + /// The elements of [`date_list`] as an array. + fn date_elements() -> ArrayRef { + let date = Scalar::extension::(TimeUnit::Days, Scalar::from(Some(42i32))); + let DType::Extension(ext_dtype) = date.dtype().clone() else { + vortex_panic!("a date is an extension value"); + }; + ExtensionArray::new( + ext_dtype, + PrimitiveArray::from_option_iter([Some(42i32), None]).into_array(), + ) + .into_array() + } } 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 0759674a902..c75db1a919f 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -16,6 +16,7 @@ use num_traits::Zero; pub use prepared::PreparedSet; pub use prepared::PreparedSetArray; pub use prepared::PreparedSetData; +pub use prepared::PreparedSetLiteral; use prost::Message; use vortex_buffer::BitBuffer; use vortex_buffer::Buffer; @@ -35,14 +36,15 @@ use crate::arrays::Constant; use crate::arrays::ConstantArray; use crate::arrays::ListViewArray; use crate::arrays::PrimitiveArray; +use crate::arrays::ScalarFn; 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; use crate::dtype::IntegerPType; use crate::dtype::Nullability; +use crate::expr::BoundExpression; use crate::expr::Expression; use crate::expr::lit; use crate::match_each_integer_ptype; @@ -209,10 +211,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) = list_array.as_opt::() - && let Some(value_scalar) = value_array.as_constant() + if let Some(value_scalar) = value_array.as_constant() + && let Some(list) = constant_list(&list_array) { - 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()); } @@ -222,16 +224,40 @@ 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 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 { + // The needle first: it is rarely constant, and a list is costly to clone. + let Some(needle) = node.child(1).as_constant() else { return Ok(None); }; - let Some(needle) = node.child(1).as_constant() else { + let Some(list) = constant_list_node(&node.child(0)) else { return Ok(None); }; let result = compute_contains_scalar(&list, &needle, options)?; Ok(Some(node.new_constant(result))) } + /// A literal list becomes a [`PreparedSetLiteral`], so that every batch the expression is + /// applied to probes one shared set rather than preparing the list again. + fn simplify( + &self, + options: &Self::Options, + expr: &BoundExpression, + ) -> VortexResult> { + let Some(list) = expr.child(0).as_opt::() else { + return Ok(None); + }; + // A null list has no set, and answers null for every needle at execution. + if list.is_null() { + return Ok(None); + } + + let set = + PreparedSetLiteral.try_new_bound_expr(PreparedSetData::try_new(list.clone())?, [])?; + Ok(Some(ListContains.try_new_bound_expr( + *options, + [set, expr.child(1).clone()], + )?)) + } + /// 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 @@ -246,7 +272,12 @@ impl ScalarFnVTable for ListContains { options: &Self::Options, expression: &Expression, ) -> VortexResult> { - let Some(list) = expression.child(0).as_opt::() else { + let list_child = expression.child(0); + let Some(list) = list_child.as_opt::().or_else(|| { + list_child + .as_opt::() + .map(PreparedSetData::list) + }) else { return Ok(None); }; if !matches!(list.dtype(), DType::List(..)) { @@ -340,35 +371,63 @@ fn compute_list_contains( .into_array()); } - if let Some(set) = array.as_opt::() { + // A literal list arrives prepared once per expression. A constant list of one batch is + // prepared here, and a prepared set that came back from the executor probes at once. + let set = if let Some(set) = array.as_opt::() { + return set.contains(value, options, ctx); + } else if let Some(set) = prepared_literal(array) { + set.clone() + } else if let Some(list) = array.as_opt::() { + PreparedSetData::try_new(list.scalar().clone())? + } else { + 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); + } + return lists_contain_needles(array, value, nullability, options, ctx); + }; + + // A canonical needle has no encoding for a kernel to use, so probe it now. + if value.is_canonical() { return set.contains(value, options, ctx); } - let nullability = options.result_nullability(array.dtype(), value.dtype()); + // 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. + let set = PreparedSetArray::new(set, array.len()).into_array(); + Ok(ListContains::try_new_opts(set, value.clone(), *options)?.into_array()) +} - if let Some(value_scalar) = value.as_constant() { - return list_contains_scalar(array, &value_scalar, nullability, options, ctx); +/// The list that every row of `array` holds: a constant, or a set prepared from one. +fn constant_list(array: &ArrayRef) -> Option<&Scalar> { + if let Some(constant) = array.as_opt::() { + return Some(constant.data().scalar()); } + if let Some(set) = array.as_opt::() { + return Some(set.data().list()); + } + prepared_literal(array).map(PreparedSetData::list) +} - if let Some(list) = array.as_opt::() { - 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() { - return set.data().contains(value, options, ctx); - } +/// The set of a [`PreparedSetLiteral`] array. +fn prepared_literal(array: &ArrayRef) -> Option<&PreparedSetData> { + array + .as_opt::()? + .data() + .scalar_fn() + .as_opt::() +} - // 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(), - ); +/// The list a reduce node holds when it is a constant or a prepared set, in an expression or an +/// array tree. +fn constant_list_node(node: &T) -> Option { + if let Some(list) = node.as_constant() { + return Some(list); } - - lists_contain_needles(array, value, nullability, options, ctx) + node.scalar_fn()? + .as_opt::() + .map(|set| set.list().clone()) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -804,6 +863,8 @@ mod tests { 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::PreparedSetData; + use crate::scalar_fn::fns::list_contains::PreparedSetLiteral; use crate::scalar_fn::fns::literal::Literal; use crate::stats::StatsSession; use crate::stats::stat as stat_expr; @@ -2036,4 +2097,46 @@ mod tests { ); Ok(()) } + + #[test] + fn optimize_prepares_a_literal_list_once() -> VortexResult<()> { + // The literal becomes a prepared set that every batch the expression is applied to + // probes, and a prepared set is left alone, so optimization reaches a fixed point. + let dtype = DType::Primitive(I32, Nullability::Nullable); + let set = i32_set(vec![Some(2), None]); + let expr = in_list(root(), lit(set.clone())).bind(&dtype)?.optimize()?; + assert_eq!( + expr.child(0) + .as_opt::() + .map(PreparedSetData::list), + Some(&set) + ); + assert_eq!(expr.optimize()?, expr); + + let mut ctx = array_session().create_execution_ctx(); + let first = PrimitiveArray::from_option_iter([Some(1i32), Some(2), None]) + .into_array() + .apply_bound(&expr)?; + assert_arrays_eq!( + first, + BoolArray::from_iter([None, Some(true), None]), + &mut ctx + ); + let second = PrimitiveArray::from_option_iter([Some(2i32), Some(3)]) + .into_array() + .apply_bound(&expr)?; + assert_arrays_eq!(second, BoolArray::from_iter([Some(true), None]), &mut ctx); + Ok(()) + } + + #[test] + fn optimize_leaves_a_null_list_a_literal() -> VortexResult<()> { + let dtype = DType::Primitive(I32, Nullability::Nullable); + let null_list = Scalar::null(DType::List(Arc::new(dtype.clone()), Nullability::Nullable)); + let expr = list_contains(lit(null_list.clone()), root()) + .bind(&dtype)? + .optimize()?; + assert_eq!(expr.child(0).as_opt::(), Some(&null_list)); + 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 index d30e2a09380..fd2933d1430 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 @@ -8,9 +8,10 @@ use std::hash::Hash; use std::hash::Hasher; use std::ops::Range; use std::sync::Arc; +use std::sync::OnceLock; use vortex_buffer::BitBuffer; -use vortex_error::VortexExpect; +use vortex_error::SharedVortexResult; use vortex_error::VortexResult; use vortex_error::vortex_bail; use vortex_error::vortex_ensure; @@ -38,15 +39,13 @@ use crate::array::ValidityVTable; use crate::array::with_empty_buffers; use crate::arrays::BoolArray; use crate::arrays::ConstantArray; -use crate::arrays::ListViewArray; +use crate::arrays::constant::list_scalar_elements; use crate::arrays::filter::FilterReduce; use crate::arrays::filter::FilterReduceAdaptor; use crate::arrays::slice::SliceReduce; use crate::arrays::slice::SliceReduceAdaptor; use crate::buffer::BufferHandle; use crate::dtype::DType; -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::ListContainsOptions; @@ -58,79 +57,131 @@ 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`]. +/// Every row holds the same list, as in a [`ConstantArray`]. [`ListContains`] gets one from a +/// [`PreparedSetLiteral`], or prepares a constant list itself, and when the needle 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. This encoding +/// The set 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. /// +/// [`PreparedSetLiteral`]: crate::scalar_fn::fns::list_contains::PreparedSetLiteral +/// /// [`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 elements of the list, and the probe built from them. +/// The data of a [`PreparedSetArray`]: the list that every row holds, and its set. +/// +/// The set is shared by every clone, so a slice or a filter of the array, and every batch a +/// [`PreparedSetLiteral`] is applied to, probe the one set. The probe needs an [`ExecutionCtx`] to +/// materialize the elements, so it is built on the first probe. +/// +/// [`PreparedSetLiteral`]: crate::scalar_fn::fns::list_contains::PreparedSetLiteral #[derive(Clone)] -pub struct PreparedSetData { - /// 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, -} +pub struct PreparedSetData(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: Box, +struct SharedSet { + /// The non-null list that every row holds. + list: Scalar, /// 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, + /// The elements, and the probe over the non-null ones, built on the first probe. + set: OnceLock>, } -impl PreparedSetData { - /// 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(); +/// The elements of the list, and the non-null ones in a probe structure. +pub(super) struct ElementSet { + /// All elements of the list, null elements included, with the element dtype of the list. + elements: ArrayRef, + pub(super) probe: Box, +} + +impl ElementSet { + fn try_new(list: &Scalar, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); // 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 probe_elements = if has_null_element { - elements.filter(valid)? - } else { + let probe_elements = if valid.all_true() { elements.clone() + } else { + elements.filter(valid)? }; - let probe = new_probe(probe_elements, ctx)?; - Ok(Self { - elements, - nullability, - set: Arc::new(ElementSet { - probe, - has_null_element, - is_empty, - }), - }) + Ok(Self { elements, probe }) } +} - /// The elements of the list that every row holds. - pub fn elements(&self) -> &ArrayRef { - &self.elements +impl PreparedSetData { + /// Prepares the non-null list scalar `list` as a set. + /// + /// # Errors + /// + /// Fails when `list` is not a list, or is null. + pub fn try_new(list: Scalar) -> VortexResult { + vortex_ensure!( + matches!(list.dtype(), DType::List(..)), + "A prepared set needs a list, got {}", + list.dtype() + ); + let (has_null_element, is_empty) = { + let Some(elements) = list.as_list().element_values() else { + vortex_bail!("A prepared set needs a non-null list"); + }; + (elements.iter().any(Option::is_none), elements.is_empty()) + }; + + Ok(Self(Arc::new(SharedSet { + list, + has_null_element, + is_empty, + set: OnceLock::new(), + }))) + } + + /// The list that every row holds. + pub fn list(&self) -> &Scalar { + &self.0.list + } + + /// The elements of the list, one row each, null elements included. + /// + /// # Errors + /// + /// Fails when the elements cannot be executed into the probe. + pub fn elements(&self, ctx: &mut ExecutionCtx) -> VortexResult<&ArrayRef> { + Ok(&self.element_set(ctx)?.elements) + } + + /// The elements and their probe, built on the first call and shared after it. + pub(super) fn element_set<'a>( + &'a self, + ctx: &mut ExecutionCtx, + ) -> VortexResult<&'a ElementSet> { + self.0 + .set + .get_or_init(|| ElementSet::try_new(&self.0.list, ctx).map_err(Arc::new)) + .as_ref() + .map_err(|err| Arc::clone(err).into()) } /// The dtype of the list that every row holds. - fn list_dtype(&self) -> DType { - DType::List(Arc::new(self.elements.dtype().clone()), self.nullability) + fn list_dtype(&self) -> &DType { + self.0.list.dtype() + } + + /// The dtype of the list's elements. + fn element_dtype(&self) -> &DType { + let DType::List(element_dtype, _) = self.list_dtype() else { + vortex_panic!("A prepared set always holds a list"); + }; + element_dtype } /// Whether each of `needles` is an element of the list, under `options`. @@ -158,7 +209,7 @@ impl PreparedSetData { } self.check_needle_dtype(needles.dtype())?; - let (bits, needle_validity) = self.set.probe.contains(needles, ctx)?; + let (bits, needle_validity) = self.element_set(ctx)?.probe.contains(needles, ctx)?; self.result_from_bits(bits, needle_validity, needles.dtype(), options) } @@ -178,7 +229,7 @@ impl PreparedSetData { ) -> 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)?; + let (bits, needle_validity) = self.element_set(ctx)?.probe.contains(&needles, ctx)?; self.result_from_bits(bits, needle_validity, needle.dtype(), options)? .execute_scalar(0, ctx) } @@ -210,11 +261,11 @@ impl PreparedSetData { ); } - 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 { + let validity = if self.0.is_empty && !options.sql_null_semantics { Validity::NonNullable - } else if options.sql_null_semantics && self.set.has_null_element { + } else if options.sql_null_semantics && self.0.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 { @@ -226,7 +277,7 @@ impl PreparedSetData { /// 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(); + let element_dtype = self.element_dtype(); if !element_dtype.eq_ignore_nullability(needle_dtype) { vortex_bail!( "Element type {} of list does not match search value {}", @@ -241,34 +292,45 @@ impl PreparedSetData { impl Debug for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { f.debug_struct("PreparedSetData") - .field("elements", &self.elements) - .field("nullability", &self.nullability) + .field("list", &self.0.list) .finish_non_exhaustive() } } impl Display for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - write!(f, "elements: {}", self.elements.len()) + write!(f, "{}", self.0.list) + } +} + +impl PartialEq for PreparedSetData { + fn eq(&self, other: &Self) -> bool { + self.0.list == other.0.list + } +} + +impl Eq for PreparedSetData {} + +impl Hash for PreparedSetData { + fn hash(&self, state: &mut H) { + self.0.list.hash(state); } } impl ArrayHash for PreparedSetData { - fn array_hash(&self, state: &mut H, accuracy: EqMode) { - self.elements.array_hash(state, accuracy); - self.nullability.hash(state); + fn array_hash(&self, state: &mut H, _accuracy: EqMode) { + self.hash(state); } } impl ArrayEq for PreparedSetData { - fn array_eq(&self, other: &Self, accuracy: EqMode) -> bool { - self.nullability == other.nullability && self.elements.array_eq(&other.elements, accuracy) + fn array_eq(&self, other: &Self, _accuracy: EqMode) -> bool { + self == other } } impl Array { - /// Prepares `elements`, the elements of a non-null list with the list nullability - /// `nullability`, as a set repeated `len` times. + /// Prepares the non-null list scalar `list` as a set repeated `len` times. /// /// [`ListContains`] prepares a constant list itself. A kernel crate can use this to test its /// [`ListContainsElementKernel`] against a prepared set. @@ -278,23 +340,17 @@ impl Array { /// /// # Errors /// - /// Fails when the elements cannot be executed into the probe. - pub 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)) + /// Fails when `list` is not a list, or is null. + pub fn try_new(list: Scalar, len: usize) -> VortexResult { + Ok(Self::new(PreparedSetData::try_new(list)?, 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(); + /// An array of `len` rows that share the prepared set `set`. + pub fn new(set: PreparedSetData, len: usize) -> Self { + let dtype = set.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)) } + unsafe { Array::from_parts_unchecked(ArrayParts::new(PreparedSet, dtype, len, set)) } } } @@ -322,7 +378,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(()) @@ -373,31 +429,10 @@ 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. - 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)) + Ok(ExecutionResult::done(ConstantArray::new( + array.data().list().clone(), + array.len(), + ))) } fn reduce_parent( @@ -415,17 +450,9 @@ impl OperationsVTable for PreparedSet { fn scalar_at( array: ArrayView<'_, PreparedSet>, _index: usize, - ctx: &mut ExecutionCtx, + _ctx: &mut ExecutionCtx, ) -> VortexResult { - 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, - )) + Ok(array.data().list().clone()) } } @@ -439,7 +466,7 @@ impl ValidityVTable for PreparedSet { impl SliceReduce for PreparedSet { fn slice(array: ArrayView<'_, Self>, range: Range) -> VortexResult> { Ok(Some( - PreparedSetArray::from_data(array.data().clone(), range.len()).into_array(), + PreparedSetArray::new(array.data().clone(), range.len()).into_array(), )) } } @@ -447,7 +474,7 @@ impl SliceReduce for PreparedSet { 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(), + PreparedSetArray::new(array.data().clone(), mask.true_count()).into_array(), )) } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/literal.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/literal.rs new file mode 100644 index 00000000000..6308bec4c5c --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/literal.rs @@ -0,0 +1,104 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use std::fmt::Formatter; + +use prost::Message; +use vortex_error::VortexResult; +use vortex_session::VortexSession; +use vortex_session::registry::CachedId; + +use super::PreparedSetArray; +use super::PreparedSetData; +use crate::ArrayRef; +use crate::ExecutionCtx; +use crate::IntoArray; +use crate::dtype::DType; +use crate::expr::Expression; +use crate::expr::display::ExprDisplay; +use crate::expr::lit; +use crate::proto::scalar as pb; +use crate::scalar::Scalar; +use crate::scalar_fn::Arity; +use crate::scalar_fn::ChildName; +use crate::scalar_fn::ExecutionArgs; +use crate::scalar_fn::ScalarFnId; +use crate::scalar_fn::ScalarFnVTable; + +/// A literal list prepared as a set for membership probes. +/// +/// [`ListContains`] puts this in place of a literal list when an expression is optimized, so that +/// every batch the expression is applied to probes one shared [`PreparedSetData`]: the list is +/// materialized, sorted and hashed once per expression rather than once per batch. Applied to an +/// array, it gives a [`PreparedSetArray`] of the batch's length in constant time. +/// +/// It serializes as the list, and deserializes into a set prepared again. +/// +/// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains +#[derive(Clone, Debug)] +pub struct PreparedSetLiteral; + +impl ScalarFnVTable for PreparedSetLiteral { + type Options = PreparedSetData; + + fn id(&self) -> ScalarFnId { + static ID: CachedId = CachedId::new("vortex.list.prepared_set_literal"); + *ID + } + + fn serialize(&self, set: &Self::Options) -> VortexResult>> { + Ok(Some(pb::Scalar::from(set.list()).encode_to_vec())) + } + + fn deserialize(&self, metadata: &[u8], session: &VortexSession) -> VortexResult { + let list = Scalar::from_proto(&pb::Scalar::decode(metadata)?, session)?; + PreparedSetData::try_new(list) + } + + fn arity(&self, _options: &Self::Options) -> Arity { + Arity::Exact(0) + } + + fn child_name(&self, _options: &Self::Options, _child_idx: usize) -> ChildName { + unreachable!() + } + + fn fmt_sql( + &self, + set: &Self::Options, + _expr: &dyn ExprDisplay, + f: &mut Formatter<'_>, + ) -> std::fmt::Result { + write!(f, "{}", set.list()) + } + + fn return_dtype(&self, set: &Self::Options, _arg_dtypes: &[DType]) -> VortexResult { + Ok(set.list().dtype().clone()) + } + + fn execute( + &self, + set: &Self::Options, + args: &dyn ExecutionArgs, + _ctx: &mut ExecutionCtx, + ) -> VortexResult { + Ok(PreparedSetArray::new(set.clone(), args.row_count()).into_array()) + } + + fn validity( + &self, + _set: &Self::Options, + _expression: &Expression, + ) -> VortexResult> { + // The list is never null. + Ok(Some(lit(true))) + } + + fn is_strict(&self, _options: &Self::Options) -> bool { + true + } + + fn is_infallible(&self, _options: &Self::Options) -> bool { + true + } +} 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 0d3028c8e01..ecae217aff4 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 @@ -5,12 +5,14 @@ //! built from its elements, so that the kernels of a needle encoding can probe their own values. mod array; +mod literal; use std::hash::BuildHasher; pub use array::PreparedSet; pub use array::PreparedSetArray; pub use array::PreparedSetData; +pub use literal::PreparedSetLiteral; use num_traits::ToPrimitive; use num_traits::WrappingSub; use vortex_buffer::BitBuffer; @@ -233,8 +235,10 @@ impl BytesSet { let buffers = data_buffers(&elements); let mut heads = HeadFilter::with_capacity(views.len()); - let mut short = HashTable::new(); - let mut long = HashTable::new(); + // Sized up front, so that no insert grows a table and hashes every element again. + let short_len = views.iter().filter(|view| view.is_inlined()).count(); + let mut short = HashTable::with_capacity(short_len); + let mut long = HashTable::with_capacity(views.len() - short_len); for (idx, view) in views.iter().enumerate() { heads.insert(view_head(view)); 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 c5ca8c3b174..e514648ccfb 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 @@ -22,7 +22,6 @@ 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; use crate::dtype::DType; @@ -39,16 +38,26 @@ 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 non-null list scalar `list`. +fn prepare(list: &Scalar) -> VortexResult { + PreparedSetData::try_new(list.clone()) } -/// 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()) +/// Prepares the non-null list scalar `list` as a set repeated `len` times. +fn prepare_array(list: &Scalar, len: usize) -> VortexResult { + Ok(PreparedSetArray::try_new(list.clone(), len)?.into_array()) +} + +/// A non-null list scalar of the rows of `elements`. +fn list_scalar(elements: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult { + let values = (0..elements.len()) + .map(|idx| elements.execute_scalar(idx, ctx)) + .collect::>>()?; + Ok(Scalar::list( + elements.dtype().clone(), + values, + Nullability::NonNullable, + )) } fn nested_needles() -> ArrayRef { @@ -116,7 +125,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 = prepare(&list, &mut ctx)?; + let set = prepare(&list)?; let result = set.contains(&needles, &options, &mut ctx)?; // The old fallback returned a lazy OR tree here. assert!(result.is::()); @@ -183,8 +192,8 @@ fn test_decimal_bitmap_across_storage_widths( .into_array() }); let options = ListContainsOptions { sql_null_semantics }; - let set = prepare(&list, &mut ctx)?; - assert!(set.set.probe.is_bitmap()); + let set = prepare(&list)?; + assert!(set.element_set(&mut ctx)?.probe.is_bitmap()); let non_match = (!sql_null_semantics).then_some(false); assert_arrays_eq!( set.contains(&needles, &options, &mut ctx)?, @@ -227,8 +236,8 @@ fn test_decimal_wide_values( decimal, ) .into_array(); - let set = prepare(&list, &mut ctx)?; - assert_eq!(set.set.probe.is_bitmap(), dense); + let set = prepare(&list)?; + assert_eq!(set.element_set(&mut ctx)?.probe.is_bitmap(), dense); assert_arrays_eq!( set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, BoolArray::from_iter([ @@ -272,8 +281,8 @@ fn test_primitive_integers( #[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); + let set = PreparedSetData::try_new(list_scalar(&elements, &mut ctx)?)?; + assert_eq!(set.element_set(&mut ctx)?.probe.is_bitmap(), dense); assert_arrays_eq!( set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, BoolArray::from_iter(expected), @@ -298,7 +307,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 = prepare_array(&set_with_null(), 4, &mut ctx)?; + let set = prepare_array(&set_with_null(), 4)?; assert_arrays_eq!(set, ConstantArray::new(set_with_null(), 4), &mut ctx); @@ -320,7 +329,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 = prepare_array(&set_with_null(), 4, &mut ctx)?; + let set = prepare_array(&set_with_null(), 4)?; let needles = PrimitiveArray::from_option_iter([Some(1i32), Some(2), None, Some(3)]).into_array(); @@ -338,7 +347,7 @@ fn test_constant_needle_gives_a_constant( ) -> VortexResult<()> { // `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 = prepare_array(&set_with_null(), 4, &mut ctx)?; + let set = prepare_array(&set_with_null(), 4)?; let needle = ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 4).into_array(); let result = ListContains::try_new_opts(set, needle, options)? @@ -361,7 +370,7 @@ fn test_constant_row_needle_probes_one_row() -> VortexResult<()> { 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 set = prepare_array(&list, 4)?; let needle = ConstantArray::new(member, 4).into_array(); let result = ListContains::try_new_opts(set, needle, ListContainsOptions::default())? .into_array() @@ -382,7 +391,7 @@ fn test_result_from_bits_applies_null_semantics( ) -> 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 set = prepare(&set_with_null())?; 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]); @@ -394,15 +403,16 @@ fn test_result_from_bits_applies_null_semantics( #[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 set = prepare(&set_with_null())?; let options = ListContainsOptions::default(); let bits = BitBuffer::from_iter([true, false]); - let short_validity = Validity::from_iter([true]); + // A validity with a null keeps its row count; an all-valid one folds to `AllValid`, which + // has none to check. + let long_validity = Validity::from_iter([true, false, true]); let needle_dtype = DType::Primitive(PType::I32, Nullability::Nullable); assert!( - set.result_from_bits(bits.clone(), short_validity, &needle_dtype, &options) + set.result_from_bits(bits.clone(), long_validity, &needle_dtype, &options) .is_err() ); @@ -426,7 +436,7 @@ fn test_bytes_set_short_and_long_views() -> VortexResult<()> { Some("ab"), ]) .into_array(); - let set = PreparedSetData::try_new(elements, Nullability::NonNullable, &mut ctx)?; + let set = PreparedSetData::try_new(list_scalar(&elements, &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. diff --git a/vortex-array/src/scalar_fn/session.rs b/vortex-array/src/scalar_fn/session.rs index 4e85b4d078c..6e5d5f20014 100644 --- a/vortex-array/src/scalar_fn/session.rs +++ b/vortex-array/src/scalar_fn/session.rs @@ -24,6 +24,7 @@ use crate::scalar_fn::fns::is_not_null::IsNotNull; use crate::scalar_fn::fns::is_null::IsNull; use crate::scalar_fn::fns::like::Like; use crate::scalar_fn::fns::list_contains::ListContains; +use crate::scalar_fn::fns::list_contains::PreparedSetLiteral; use crate::scalar_fn::fns::list_length::ListLength; use crate::scalar_fn::fns::list_sum::ListSum; use crate::scalar_fn::fns::literal::Literal; @@ -76,6 +77,7 @@ impl Default for ScalarFnSession { this.register(IsNull); this.register(Like); this.register(ListContains); + this.register(PreparedSetLiteral); this.register(ListLength); this.register(ListSum); this.register(Literal); From 6791e82d32e23e2247873ece7f9838e1313f467f Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Tue, 29 Sep 2026 15:48:17 -0700 Subject: [PATCH 09/10] more Signed-off-by: Robert Kruszewski --- .../src/scalar_fn/fns/list_contains/mod.rs | 12 +++-- .../fns/list_contains/prepared/array.rs | 47 +++++++++++++------ 2 files changed, 40 insertions(+), 19 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 c75db1a919f..0e78b19efe8 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -242,16 +242,20 @@ impl ScalarFnVTable for ListContains { options: &Self::Options, expr: &BoundExpression, ) -> VortexResult> { - let Some(list) = expr.child(0).as_opt::() else { + let Some(literal) = expr.child(0).as_scalar() else { return Ok(None); }; // A null list has no set, and answers null for every needle at execution. - if list.is_null() { + if !literal + .as_opt::() + .is_some_and(|list| !list.is_null()) + { return Ok(None); } - let set = - PreparedSetLiteral.try_new_bound_expr(PreparedSetData::try_new(list.clone())?, [])?; + // The set shares the literal, so the list is not cloned. + let set = PreparedSetLiteral + .try_new_bound_expr(PreparedSetData::try_from_literal(literal.clone())?, [])?; Ok(Some(ListContains.try_new_bound_expr( *options, [set, expr.child(1).clone()], 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 fd2933d1430..4cf904a342a 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 @@ -48,7 +48,10 @@ use crate::buffer::BufferHandle; use crate::dtype::DType; use crate::optimizer::rules::ParentRuleSet; use crate::scalar::Scalar; +use crate::scalar_fn::ScalarFnRef; +use crate::scalar_fn::ScalarFnVTableExt; use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::scalar_fn::fns::literal::Literal; use crate::serde::ArrayChildren; use crate::validity::Validity; @@ -84,8 +87,9 @@ pub struct PreparedSet; pub struct PreparedSetData(Arc); struct SharedSet { - /// The non-null list that every row holds. - list: Scalar, + /// The [`Literal`] of the non-null list that every row holds, shared with the expression it + /// came from rather than cloned out of it. + literal: ScalarFnRef, /// 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. @@ -125,12 +129,25 @@ impl PreparedSetData { /// /// Fails when `list` is not a list, or is null. pub fn try_new(list: Scalar) -> VortexResult { - vortex_ensure!( - matches!(list.dtype(), DType::List(..)), - "A prepared set needs a list, got {}", - list.dtype() - ); + Self::try_from_literal(Literal.bind(list)) + } + + /// Prepares the non-null list that the [`Literal`] `literal` holds as a set, sharing the + /// literal rather than cloning its list. + /// + /// # Errors + /// + /// Fails when `literal` is not a literal, or holds a null, or a value that is not a list. + pub(crate) fn try_from_literal(literal: ScalarFnRef) -> VortexResult { let (has_null_element, is_empty) = { + let Some(list) = literal.as_opt::() else { + vortex_bail!("A prepared set needs a literal list"); + }; + vortex_ensure!( + matches!(list.dtype(), DType::List(..)), + "A prepared set needs a list, got {}", + list.dtype() + ); let Some(elements) = list.as_list().element_values() else { vortex_bail!("A prepared set needs a non-null list"); }; @@ -138,7 +155,7 @@ impl PreparedSetData { }; Ok(Self(Arc::new(SharedSet { - list, + literal, has_null_element, is_empty, set: OnceLock::new(), @@ -147,7 +164,7 @@ impl PreparedSetData { /// The list that every row holds. pub fn list(&self) -> &Scalar { - &self.0.list + self.0.literal.as_::() } /// The elements of the list, one row each, null elements included. @@ -166,14 +183,14 @@ impl PreparedSetData { ) -> VortexResult<&'a ElementSet> { self.0 .set - .get_or_init(|| ElementSet::try_new(&self.0.list, ctx).map_err(Arc::new)) + .get_or_init(|| ElementSet::try_new(self.list(), ctx).map_err(Arc::new)) .as_ref() .map_err(|err| Arc::clone(err).into()) } /// The dtype of the list that every row holds. fn list_dtype(&self) -> &DType { - self.0.list.dtype() + self.list().dtype() } /// The dtype of the list's elements. @@ -292,20 +309,20 @@ impl PreparedSetData { impl Debug for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { f.debug_struct("PreparedSetData") - .field("list", &self.0.list) + .field("list", self.list()) .finish_non_exhaustive() } } impl Display for PreparedSetData { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0.list) + write!(f, "{}", self.list()) } } impl PartialEq for PreparedSetData { fn eq(&self, other: &Self) -> bool { - self.0.list == other.0.list + self.list() == other.list() } } @@ -313,7 +330,7 @@ impl Eq for PreparedSetData {} impl Hash for PreparedSetData { fn hash(&self, state: &mut H) { - self.0.list.hash(state); + self.list().hash(state); } } From 77467d3b6762405901dd681ae51cdf03666465b6 Mon Sep 17 00:00:00 2001 From: Robert Kruszewski Date: Tue, 29 Sep 2026 16:33:35 -0700 Subject: [PATCH 10/10] move Signed-off-by: Robert Kruszewski --- .../src/scalar_fn/fns/list_contains/mod.rs | 44 +++++++++---------- 1 file changed, 20 insertions(+), 24 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 0e78b19efe8..3dd720353fe 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -44,7 +44,6 @@ use crate::arrays::primitive::PrimitiveArrayExt; use crate::dtype::DType; use crate::dtype::IntegerPType; use crate::dtype::Nullability; -use crate::expr::BoundExpression; use crate::expr::Expression; use crate::expr::lit; use crate::match_each_integer_ptype; @@ -222,27 +221,24 @@ 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 looks a constant needle up at execution. + /// array trees alike. Otherwise a literal list becomes a [`PreparedSetLiteral`], so that every + /// batch the expression is applied to probes one shared set rather than preparing the list + /// again. fn reduce(&self, options: &Self::Options, node: &T) -> VortexResult> { + let list = node.child(0); + // The needle first: it is rarely constant, and a list is costly to clone. - let Some(needle) = node.child(1).as_constant() else { - return Ok(None); - }; - let Some(list) = constant_list_node(&node.child(0)) else { - return Ok(None); - }; - let result = compute_contains_scalar(&list, &needle, options)?; - Ok(Some(node.new_constant(result))) - } + if let Some(needle) = node.child(1).as_constant() + && let Some(list) = constant_list_node(&list) + { + let result = compute_contains_scalar(&list, &needle, options)?; + return Ok(Some(node.new_constant(result))); + } - /// A literal list becomes a [`PreparedSetLiteral`], so that every batch the expression is - /// applied to probes one shared set rather than preparing the list again. - fn simplify( - &self, - options: &Self::Options, - expr: &BoundExpression, - ) -> VortexResult> { - let Some(literal) = expr.child(0).as_scalar() else { + // Only an expression tree holds a literal as a scalar function: applied to an array, a + // literal is a constant array. So this prepares the list once per expression, and never + // once per batch. + let Some(literal) = list.scalar_fn() else { return Ok(None); }; // A null list has no set, and answers null for every needle at execution. @@ -254,11 +250,11 @@ impl ScalarFnVTable for ListContains { } // The set shares the literal, so the list is not cloned. - let set = PreparedSetLiteral - .try_new_bound_expr(PreparedSetData::try_from_literal(literal.clone())?, [])?; - Ok(Some(ListContains.try_new_bound_expr( - *options, - [set, expr.child(1).clone()], + let set = PreparedSetData::try_from_literal(literal.clone())?; + let set = node.new_node(PreparedSetLiteral.bind(set), &[])?; + Ok(Some(node.new_node( + ListContains.bind(*options), + &[set, node.child(1)], )?)) }