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/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/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..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; @@ -43,7 +47,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 +279,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 +310,124 @@ 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. +pub(crate) fn list_scalar_elements(list: &ListScalar, allocator: &BufferAllocatorRef) -> ArrayRef { + let element_dtype = list.element_dtype(); + let Some(elements) = list.element_values() else { + return Canonical::empty(element_dtype).into_array(); + }; + flat_elements(element_dtype, elements, allocator) +} + +/// 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"), + } + } + 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))) + } + } +} + /// Creates a [`FixedSizeListArray`] whose every row holds the same list. fn constant_canonical_fixed_size_list_array( values: Option>, @@ -405,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; @@ -428,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; @@ -1021,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/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..952ac16d8c9 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs @@ -20,11 +20,8 @@ use crate::scalar_fn::fns::list_contains::ListContainsOptions; /// 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 { @@ -40,6 +37,19 @@ 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. 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/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs index f67536a960b..3dd720353fe 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,30 @@ // SPDX-FileCopyrightText: Copyright the Vortex contributors mod kernel; +mod prepared; use std::fmt::Display; use std::fmt::Formatter; +use std::hash::Hash; +use std::iter; use std::ops::BitOr; 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; +pub use prepared::PreparedSetLiteral; 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; @@ -29,26 +36,29 @@ 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::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; 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; 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 +209,93 @@ 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() - && let Some(value_scalar) = value_array.as_constant() + // Borrow the list: a constant list scalar owns every element, so cloning it is not free. + 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()); } compute_list_contains(&list_array, &value_array, options, ctx) } + /// A constant list and a constant needle fold to their constant answer, in expression and + /// 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. + 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))); + } + + // 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. + if !literal + .as_opt::() + .is_some_and(|list| !list.is_null()) + { + return Ok(None); + } + + // The set shares the literal, so the list is not cloned. + 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)], + )?)) + } + + /// 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, + expression: &Expression, + ) -> VortexResult> { + 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(..)) { + 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 // needle is null against any list, as a null list always is. fn is_strict(&self, options: &Self::Options) -> bool { @@ -284,70 +371,63 @@ fn compute_list_contains( .into_array()); } - 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); - } + // 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); + }; - if let Some(list_scalar) = array.as_constant() { - return constant_list_scalar_contains(&list_scalar.as_list(), value, nullability, options); + // 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); } - todo!("unsupported list contains with list and element as arrays") + // 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()) } -/// 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()?)?; +/// 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 result.dtype().nullability() != nullability { - result = result.cast(DType::Bool(nullability))?; + if let Some(set) = array.as_opt::() { + return Some(set.data().list()); } + prepared_literal(array).map(PreparedSetData::list) +} - Ok(result) +/// The set of a [`PreparedSetLiteral`] array. +fn prepared_literal(array: &ArrayRef) -> Option<&PreparedSetData> { + array + .as_opt::()? + .data() + .scalar_fn() + .as_opt::() +} + +/// 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); + } + node.scalar_fn()? + .as_opt::() + .map(|set| set.list().clone()) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -439,6 +519,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 +814,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 +829,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; @@ -623,6 +850,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; @@ -630,9 +858,14 @@ 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; 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; use crate::validity::Validity; @@ -1262,6 +1495,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 +1527,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 +1655,158 @@ 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])) + } + + /// 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 +1834,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. @@ -1480,4 +2038,105 @@ 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(()) + } + + #[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 new file mode 100644 index 00000000000..4cf904a342a --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -0,0 +1,497 @@ +// 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 std::sync::OnceLock; + +use vortex_buffer::BitBuffer; +use vortex_error::SharedVortexResult; +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 super::new_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::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::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; + +/// 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`]. [`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 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 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(Arc); + +struct SharedSet { + /// 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. + is_empty: bool, + /// The elements, and the probe over the non-null ones, built on the first probe. + set: OnceLock>, +} + +/// 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 probe_elements = if valid.all_true() { + elements.clone() + } else { + elements.filter(valid)? + }; + let probe = new_probe(probe_elements, ctx)?; + + Ok(Self { elements, probe }) + } +} + +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 { + 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"); + }; + (elements.iter().any(Option::is_none), elements.is_empty()) + }; + + Ok(Self(Arc::new(SharedSet { + literal, + has_null_element, + is_empty, + set: OnceLock::new(), + }))) + } + + /// The list that every row holds. + pub fn list(&self) -> &Scalar { + self.0.literal.as_::() + } + + /// 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.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.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`. + /// + /// 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. + /// + /// 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 { + 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.element_set(ctx)?.probe.contains(needles, ctx)?; + 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, 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.element_set(ctx)?.probe.contains(&needles, ctx)?; + self.result_from_bits(bits, needle_validity, needle.dtype(), options)? + .execute_scalar(0, ctx) + } + + /// 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.0.is_empty && !options.sql_null_semantics { + Validity::NonNullable + } 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 { + needle_validity + }; + + 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.element_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 { + 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, "{}", self.list()) + } +} + +impl PartialEq for PreparedSetData { + fn eq(&self, other: &Self) -> bool { + self.list() == other.list() + } +} + +impl Eq for PreparedSetData {} + +impl Hash for PreparedSetData { + fn hash(&self, state: &mut H) { + self.list().hash(state); + } +} + +impl ArrayHash for PreparedSetData { + 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 == other + } +} + +impl Array { + /// 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. + /// + /// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains + /// [`ListContainsElementKernel`]: crate::scalar_fn::fns::list_contains::ListContainsElementKernel + /// + /// # Errors + /// + /// 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 `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, set)) } + } +} + +const PARENT_RULES: ParentRuleSet = ParentRuleSet::new(&[ + ParentRuleSet::lift(&FilterReduceAdaptor(PreparedSet)), + ParentRuleSet::lift(&SliceReduceAdaptor(PreparedSet)), +]); + +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.data().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::new(array.data().clone(), range.len()).into_array(), + )) + } +} + +impl FilterReduce for PreparedSet { + fn filter(array: ArrayView<'_, Self>, mask: &Mask) -> VortexResult> { + Ok(Some( + 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 new file mode 100644 index 00000000000..ecae217aff4 --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -0,0 +1,519 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +//! 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; +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; +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_mask::Mask; +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::converted_buffer; +use crate::arrays::primitive::PrimitiveArrayExt; +use crate::arrays::varbinview::BinaryView; +use crate::dtype::DType; +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_fn::fns::binary::build_row_comparator; +use crate::scalar_fn::fns::binary::collect_bits; +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; +/// 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. +/// +/// 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)>; + + /// 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 { + 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); + +impl Probe for PrimitiveSet { + fn contains( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult<(BitBuffer, 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()?)) + } + + #[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)?; + 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 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 { + heads: HeadFilter, + hasher: RandomState, + /// 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 { + fn new(elements: VarBinViewArray) -> Self { + let hasher = RandomState::default(); + let views = elements.views(); + let buffers = data_buffers(&elements); + + let mut heads = HeadFilter::with_capacity(views.len()); + // 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)); + + 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(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 { + heads, + hasher, + 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 { + fn contains( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult<(BitBuffer, Validity)> { + 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| 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, + 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 }) + } +} + +impl Probe for RowSet { + fn contains( + &self, + needles: &ArrayRef, + ctx: &mut ExecutionCtx, + ) -> VortexResult<(BitBuffer, Validity)> { + if self.indices.is_empty() { + return Ok(( + 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(&self.elements, &needles, ctx)?; + let bits = BitBuffer::collect_bool_in( + needles.len(), + |row| { + self.indices + .binary_search_by(|&element| compare(element, row)) + .is_ok() + }, + ctx.allocator().clone(), + ); + Ok((bits, needles.validity()?)) + } +} + +/// 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 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, + } +} + +/// 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 + Send + Sync + 'static { + 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) +} + +#[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..e514648ccfb --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -0,0 +1,468 @@ +// 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::PreparedSet; +use super::PreparedSetArray; +use crate::ArrayRef; +use crate::ExecutionCtx; +use crate::IntoArray; +use crate::VortexSessionExecute; +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::arrays::VarBinViewArray; +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::ListContains; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::scalar_fn::fns::list_contains::PreparedSetData; +use crate::validity::Validity; + +/// Prepares the non-null list scalar `list`. +fn prepare(list: &Scalar) -> VortexResult { + PreparedSetData::try_new(list.clone()) +} + +/// 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 { + 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() +} + +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(struct_needles())] +#[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 = Scalar::list(dtype, elements, Nullability::NonNullable); + let options = ListContainsOptions { sql_null_semantics }; + let set = prepare(&list)?; + 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()) + .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::>>()?; + + // 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(()) +} + +#[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 = 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())), + decimal, + ) + .into_array() + }); + let options = ListContainsOptions { sql_null_semantics }; + 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)?, + 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 = 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::( + [ + Some(base), + Some(second), + Some(base - i256::ONE), + Some(i256::ZERO), + Some(base + i256::from_i128(1i128 << 64)), + None, + ], + decimal, + ) + .into_array(); + 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([ + Some(true), + Some(true), + Some(false), + Some(false), + Some(false), + None, + ]), + &mut ctx + ); + 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(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), + &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 = prepare_array(&set_with_null(), 4)?; + + 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 = prepare_array(&set_with_null(), 4)?; + 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_gives_a_constant( + #[case] options: ListContainsOptions, + #[case] expected: Option, +) -> 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)?; + let needle = ConstantArray::new(Scalar::primitive(3i32, Nullability::Nullable), 4).into_array(); + + let result = ListContains::try_new_opts(set, needle, options)? + .into_array() + .execute::(&mut ctx)?; + let expected = match expected { + Some(value) => Scalar::bool(value, Nullability::Nullable), + None => Scalar::null(DType::Bool(Nullability::Nullable)), + }; + 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)?; + 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(()) +} + +#[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())?; + 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 set = prepare(&set_with_null())?; + let options = ListContainsOptions::default(); + let bits = BitBuffer::from_iter([true, false]); + + // 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(), long_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(()) +} + +#[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(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. + 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(()) +} 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);