diff --git a/docs/api/python/expr.rst b/docs/api/python/expr.rst index 9ee4a7d4450..17717387bba 100644 --- a/docs/api/python/expr.rst +++ b/docs/api/python/expr.rst @@ -52,6 +52,7 @@ Expressions are picklable, so a filter built in one process can be sent to anoth ~vortex.expr.pack ~vortex.expr.merge ~vortex.expr.list_contains + ~vortex.expr.in_list ~vortex.expr.list_length ~vortex.expr.list_sum ~vortex.expr.case_when @@ -153,6 +154,8 @@ Lists .. autofunction:: vortex.expr.list_contains +.. autofunction:: vortex.expr.in_list + .. autofunction:: vortex.expr.list_length .. autofunction:: vortex.expr.list_sum diff --git a/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java b/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java index b2e2a8be875..2608b2a41ba 100644 --- a/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java +++ b/java/vortex-jni/src/main/java/dev/vortex/api/Expression.java @@ -154,6 +154,53 @@ public static Expression between( value.nativePointer(), lower.nativePointer(), upper.nativePointer(), lowerStrict, upperStrict)); } + /** + * Whether the list {@code list} evaluates to contains {@code needle}, with Vortex's null rules: a null element + * never matches anything, so a needle that matches no element is {@code false}. A null {@code needle} is null, + * except against an empty list, where the result is {@code false}. + * + *

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

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

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

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

The large-set case is the reason the binding exists. A caller without it has to expand {@code IN} into a chain of + * equality comparisons, which is why callers cap the set size and fall back to a bounding range; a single + * {@code list_contains} node carries the whole set, and the stats rewrite falsifies it for a zone only when every + * element misses that zone's bounds. + */ +public final class ListContainsFilterTest { + private static final int ROW_COUNT = 6; + + @TempDir + static Path tempDir; + + private static Session session; + private static String filePath; + + @BeforeAll + public static void loadLibrary() { + NativeLoader.loadJni(); + } + + @BeforeAll + static void writeFile() throws IOException { + session = Session.create(); + filePath = + tempDir.resolve("list_contains.vortex").toAbsolutePath().toUri().toString(); + + BufferAllocator allocator = ArrowAllocation.rootAllocator(); + Schema schema = new Schema(List.of( + Field.notNullable("id", new ArrowType.Int(32, true)), + Field.notNullable("name", new ArrowType.Utf8()), + Field.nullable("maybe", new ArrowType.Int(32, true)))); + + try (VortexWriter writer = VortexWriter.builder(session, filePath, schema, allocator) + .build(); + VectorSchemaRoot root = VectorSchemaRoot.create(schema, allocator)) { + IntVector id = (IntVector) root.getVector("id"); + VarCharVector name = (VarCharVector) root.getVector("name"); + IntVector maybe = (IntVector) root.getVector("maybe"); + id.allocateNew(ROW_COUNT); + name.allocateNew(ROW_COUNT); + maybe.allocateNew(ROW_COUNT); + for (int i = 0; i < ROW_COUNT; i++) { + id.setSafe(i, i + 1); + name.setSafe(i, ("row-" + (i + 1)).getBytes(UTF_8)); + // Every other row is null; the rest carry their id. + if (i % 2 == 0) { + maybe.setNull(i); + } else { + maybe.setSafe(i, i + 1); + } + } + root.setRowCount(ROW_COUNT); + + try (ArrowArray array = ArrowArray.allocateNew(allocator); + ArrowSchema arrowSchema = ArrowSchema.allocateNew(allocator)) { + Data.exportVectorSchemaRoot(allocator, root, null, array, arrowSchema); + writer.writeBatch(array.memoryAddress(), arrowSchema.memoryAddress()); + } + writer.finish(); + } + } + + @Test + public void inKeepsOnlyTheMatchingRows() { + Expression set = Expression.literalList(Expression.literal(2), Expression.literal(4), Expression.literal(99)); + assertEquals(List.of(2, 4), scanIds(Expression.in(Expression.column("id"), set))); + } + + @Test + public void notInKeepsTheComplement() { + Expression set = Expression.literalList(Expression.literal(2), Expression.literal(4), Expression.literal(99)); + assertEquals(List.of(1, 3, 5, 6), scanIds(Expression.notIn(Expression.column("id"), set))); + } + + @Test + public void listContainsTakesTheListFirst() { + // in(value, list) is listContains(list, value) with the operands the other way round; both spellings have to + // agree, and passing the column as the list would be a type error rather than a silently different filter. + Expression set = Expression.literalList(Expression.literal(3)); + assertEquals(List.of(3), scanIds(Expression.listContains(set, Expression.column("id")))); + } + + @Test + public void inWithANullElementKeepsOnlyMatches() { + // maybe = [null, 2, null, 4, null, 6]; the set holds 2 and a null. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(2), scanIds(Expression.in(Expression.column("maybe"), set))); + } + + @Test + public void notInWithANullElementKeepsNothing() { + // SQL: x NOT IN (2, NULL) is false for 2 and null for everything else, so no row survives. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(), scanIds(Expression.notIn(Expression.column("maybe"), set))); + // Without the null element NOT IN keeps the non-null non-matches. + assertEquals( + List.of(4, 6), + scanIds(Expression.notIn(Expression.column("maybe"), Expression.literalList(Expression.literal(2))))); + } + + @Test + public void vortexNullRulesTreatANullElementAsNoMatch() { + // listContains without SQL semantics: the null element never matches, so a non-match stays false and its + // negation keeps the non-null non-matches. + Expression set = Expression.literalList(Expression.literal(2), Expression.nullLiteral(Expression.DType.I32)); + assertEquals(List.of(2), scanIds(Expression.listContains(set, Expression.column("maybe")))); + assertEquals(List.of(4, 6), scanIds(Expression.not(Expression.listContains(set, Expression.column("maybe"))))); + } + + @Test + public void aStringSetFiltersOnUtf8() { + Expression set = Expression.literalList(Expression.literal("row-1"), Expression.literal("row-6")); + assertEquals(List.of(1, 6), scanIds(Expression.in(Expression.column("name"), set))); + } + + @Test + public void aSetLargerThanAnyOrChainCapStillPushesDown() { + // The whole point of the binding: 1000 literals stay one expression node. Only 5 is present in the file. + Expression[] elements = new Expression[1000]; + for (int i = 0; i < elements.length; i++) { + elements[i] = Expression.literal(i == 0 ? 5 : ROW_COUNT + i); + } + assertEquals(List.of(5), scanIds(Expression.in(Expression.column("id"), Expression.literalList(elements)))); + } + + @Test + public void anEmptySetMatchesNothing() { + Expression empty = Expression.literalEmptyList(Expression.DType.I32); + assertTrue(scanIds(Expression.in(Expression.column("id"), empty)).isEmpty()); + } + + @Test + public void aNullSetMatchesNothing() { + // A null list yields null rather than false, which filters the row out just the same. + Expression nullList = Expression.nullLiteralList(Expression.DType.I32); + assertTrue(scanIds(Expression.in(Expression.column("id"), nullList)).isEmpty()); + } + + /** Reads the {@code id} column of every row that survives {@code filter}, in file order. */ + private static List scanIds(Expression filter) { + BufferAllocator allocator = ArrowAllocation.rootAllocator(); + DataSource dataSource = DataSource.open(session, filePath); + Scan scan = dataSource.scan( + ScanOptions.builder().filter(filter).ordered(true).build()); + + List ids = new ArrayList<>(); + while (scan.hasNext()) { + Partition partition = scan.next(); + try (ArrowReader reader = partition.scanArrow(allocator)) { + while (reader.loadNextBatch()) { + VectorSchemaRoot root = reader.getVectorSchemaRoot(); + IntVector id = (IntVector) root.getVector("id"); + for (int i = 0; i < root.getRowCount(); i++) { + ids.add(id.get(i)); + } + } + } catch (IOException e) { + throw new AssertionError("failed reading partition", e); + } + } + return ids; + } +} diff --git a/vortex-array/benches/list_contains_set.rs b/vortex-array/benches/list_contains_set.rs index 87b2057f9c2..76e75f33dcf 100644 --- a/vortex-array/benches/list_contains_set.rs +++ b/vortex-array/benches/list_contains_set.rs @@ -60,7 +60,7 @@ fn random_i64(len: usize) -> (Vec, Vec) { fn bench_in_set(bencher: Bencher, set: Scalar, needles: ArrayRef) { let session = vortex_array::array_session(); - // Optimized as a scan optimizes it, so the set arrives normalized. + // Optimized as a scan optimizes it. let expr = list_contains(lit(set), root()) .bind(needles.dtype()) .unwrap() diff --git a/vortex-array/src/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..7c5ff18106f 100644 --- a/vortex-array/src/arrays/constant/vtable/canonical.rs +++ b/vortex-array/src/arrays/constant/vtable/canonical.rs @@ -43,6 +43,7 @@ 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::validity::Validity; @@ -273,25 +274,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 +305,31 @@ 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(); + }; + + 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() +} + /// Creates a [`FixedSizeListArray`] whose every row holds the same list. fn constant_canonical_fixed_size_list_array( values: Option>, diff --git a/vortex-array/src/scalar/typed_view/list.rs b/vortex-array/src/scalar/typed_view/list.rs index f97857c92ed..359a1b32fcb 100644 --- a/vortex-array/src/scalar/typed_view/list.rs +++ b/vortex-array/src/scalar/typed_view/list.rs @@ -160,6 +160,14 @@ impl<'a> ListScalar<'a> { }) } + /// Returns the values of the list's elements, `None` for a null element, without building a + /// scalar for each. + /// + /// Returns None if the list is null. + pub(crate) fn element_values(&self) -> Option<&'a [Option]> { + self.elements + } + /// Returns all elements in the list as a vector of scalars. /// /// Returns None if the list is null. diff --git a/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs b/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs index a1635c718ff..8a27d18092e 100644 --- a/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs +++ b/vortex-array/src/scalar_fn/fns/binary/compare/mod.rs @@ -48,6 +48,7 @@ mod boolean; mod bytes; mod decimal; mod nested; +pub(crate) use nested::build_comparator as build_row_comparator; mod primitive; #[cfg(test)] mod tests; @@ -304,7 +305,7 @@ pub(super) fn collect_zip_bits( } /// Bit-pack the predicate `f(values[i])` over a slice into a [`BitBuffer`]. -pub(super) fn collect_bits( +pub(crate) fn collect_bits( values: &[T], f: impl Fn(T) -> bool, allocator: &BufferAllocatorRef, diff --git a/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs b/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs index 0ce889c46f1..ccccc3ca955 100644 --- a/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs +++ b/vortex-array/src/scalar_fn/fns/binary/compare/nested.rs @@ -49,7 +49,7 @@ use crate::scalar_fn::fns::binary::compare::compare_validity; use crate::scalar_fn::fns::operators::CompareOperator; /// A row comparator: compares row `i` of the left operand against row `j` of the right operand. -type RowComparator = Box Ordering>; +pub(crate) type RowComparator = Box Ordering>; /// Compare two nested arrays row by row. pub(super) fn compare_nested( @@ -116,7 +116,7 @@ fn validity_mask(array: &ArrayRef, ctx: &mut ExecutionCtx) -> VortexResult /// Build a row comparator over two recursively canonical arrays of the same logical dtype /// (ignoring nullability). Null values order before all non-null values at every level. -fn build_comparator( +pub(crate) fn build_comparator( lhs: &ArrayRef, rhs: &ArrayRef, ctx: &mut ExecutionCtx, diff --git a/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs b/vortex-array/src/scalar_fn/fns/list_contains/kernel.rs index 7cee5f7aac1..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..0759674a902 100644 --- a/vortex-array/src/scalar_fn/fns/list_contains/mod.rs +++ b/vortex-array/src/scalar_fn/fns/list_contains/mod.rs @@ -2,23 +2,29 @@ // SPDX-FileCopyrightText: Copyright the Vortex contributors mod kernel; +mod prepared; use std::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; 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; @@ -31,24 +37,27 @@ use crate::arrays::ListViewArray; use crate::arrays::PrimitiveArray; use crate::arrays::ScalarFnArray; use crate::arrays::bool::BoolArrayExt; +use crate::arrays::constant::list_scalar_elements; use crate::arrays::listview::ListViewArraySlotsExt; use crate::arrays::primitive::PrimitiveArrayExt; -use crate::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 +208,63 @@ impl ScalarFnVTable for ListContains { let list_array = args.get(0)?; let value_array = args.get(1)?; - if let Some(list_scalar) = list_array.as_constant() + // Borrow the list: a constant list scalar owns every element, so cloning it is not free. + if let Some(list) = list_array.as_opt::() && let Some(value_scalar) = value_array.as_constant() { - let result = compute_contains_scalar(&list_scalar, &value_scalar, options)?; + let result = compute_contains_scalar(list.scalar(), &value_scalar, options)?; return Ok(ConstantArray::new(result, args.row_count()).into_array()); } compute_list_contains(&list_array, &value_array, options, ctx) } + /// A constant list and a constant needle fold to their constant answer, in expression and + /// array trees alike. A [`PreparedSetArray`] list looks a constant needle up at execution. + fn reduce(&self, options: &Self::Options, node: &T) -> VortexResult> { + let Some(list) = node.child(0).as_constant() else { + return Ok(None); + }; + let Some(needle) = node.child(1).as_constant() else { + return Ok(None); + }; + let result = compute_contains_scalar(&list, &needle, options)?; + Ok(Some(node.new_constant(result))) + } + + /// The validity of `needle IN list` for a literal list, decided without a probe. + /// + /// Off SQL null semantics only a null needle gives null, unless the list is empty, which + /// answers `false` for every needle. Under them a null element can give null too, by leaving + /// a non-match unknown, so only a list without one is decided from the needle alone. A null + /// list gives null for every needle. + /// + /// Without this the validity of a node over a constant list would execute the node, which + /// prepares the list as a set only to read the validity off the result. + fn validity( + &self, + options: &Self::Options, + expression: &Expression, + ) -> VortexResult> { + let Some(list) = expression.child(0).as_opt::() 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 +340,35 @@ fn compute_list_contains( .into_array()); } + if let Some(set) = array.as_opt::() { + return set.contains(value, options, ctx); + } + let nullability = options.result_nullability(array.dtype(), value.dtype()); if let Some(value_scalar) = value.as_constant() { return list_contains_scalar(array, &value_scalar, nullability, options, ctx); } - if let Some(list_scalar) = array.as_constant() { - return constant_list_scalar_contains(&list_scalar.as_list(), value, nullability, options); - } - - todo!("unsupported list contains with list and element as arrays") -} - -/// There is a constant list scalar (haystack) being compared to an array of needles. -/// -/// The result stays lazy. `Or` is Kleene, so under SQL null semantics the disjunction of the raw -/// comparisons is already the `IN` answer. -fn constant_list_scalar_contains( - list_scalar: &ListScalar, - values: &ArrayRef, - nullability: Nullability, - options: &ListContainsOptions, -) -> 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()); + if let Some(list) = array.as_opt::() { + let elements = list_scalar_elements(&list.scalar().as_list(), ctx.allocator()); + let set = + PreparedSetArray::try_new(elements, array.dtype().nullability(), array.len(), ctx)?; - // A 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()?)?; - } + // A canonical needle has no encoding for a kernel to use, so probe it now. + if value.is_canonical() { + return set.data().contains(value, options, ctx); + } - if result.dtype().nullability() != nullability { - result = result.cast(DType::Bool(nullability))?; + // Give the prepared set back to the executor. Thus a kernel of the needle encoding can + // probe it, for example on the values of a dictionary only. When no kernel does, the next + // execution finds the prepared set above and probes it, so this happens at most once. + return Ok( + ListContains::try_new_opts(set.into_array(), value.clone(), *options)?.into_array(), + ); } - Ok(result) + lists_contain_needles(array, value, nullability, options, ctx) } /// Returns a [`BoolArray`] where each bit represents if a list contains the scalar. @@ -439,6 +460,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 +755,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 +770,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 +791,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 +799,12 @@ 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::literal::Literal; use crate::stats::StatsSession; use crate::stats::stat as stat_expr; use crate::validity::Validity; @@ -1262,6 +1434,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 +1466,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 +1594,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 +1773,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 +1977,63 @@ mod tests { ); Ok(()) } + + #[rstest] + #[case::default(ListContainsOptions::default(), Some(false))] + #[case::sql(SQL, None)] + fn constant_needle_in_constant_set_folds_to_a_constant( + #[case] options: ListContainsOptions, + #[case] expected: Option, + ) -> VortexResult<()> { + // `3 IN (2, NULL)` is decided at optimization, in an expression and in an array tree. + let set = i32_set(vec![Some(2), None]); + let needle = Scalar::primitive(3i32, Nullability::Nullable); + let expected = match expected { + Some(value) => Scalar::bool(value, Nullability::Nullable), + None => Scalar::null(DType::Bool(Nullability::Nullable)), + }; + + let expr = list_contains_opts(lit(set.clone()), lit(needle.clone()), options) + .bind(&DType::Primitive(I32, Nullability::Nullable))? + .optimize()?; + assert_eq!(expr.as_opt::(), Some(&expected)); + + let array = ListContains::try_new_opts( + ConstantArray::new(set, 3).into_array(), + ConstantArray::new(needle, 3).into_array(), + options, + )? + .into_array() + .optimize()?; + assert_eq!(array.as_constant(), Some(expected)); + Ok(()) + } + + #[test] + fn validity_of_a_constant_set_needs_no_probe() -> VortexResult<()> { + // Off SQL null semantics the needle alone decides validity. Under them a null element + // makes validity depend on the probe, which only the fallback computes. + let needle = col("a"); + let with_null = lit(i32_set(vec![Some(2), None])); + let without_null = lit(i32_set(vec![Some(2)])); + + let expr = list_contains(with_null.clone(), needle.clone()); + assert_eq!(expr.validity()?, needle.validity()?); + let expr = in_list(needle.clone(), without_null); + assert_eq!(expr.validity()?, needle.validity()?); + let expr = in_list(needle.clone(), with_null); + assert_eq!(expr.validity()?, is_not_null(expr)); + + let expr = list_contains(lit(empty_i32_list()), needle.clone()); + assert_eq!(expr.validity()?, lit(true)); + let null_list = Scalar::null(DType::List( + Arc::new(DType::Primitive(I32, Nullability::Nullable)), + Nullability::Nullable, + )); + assert_eq!( + list_contains(lit(null_list), needle).validity()?, + lit(false) + ); + Ok(()) + } } diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs new file mode 100644 index 00000000000..d30e2a09380 --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/array.rs @@ -0,0 +1,453 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright the Vortex contributors + +use std::fmt::Debug; +use std::fmt::Display; +use std::fmt::Formatter; +use std::hash::Hash; +use std::hash::Hasher; +use std::ops::Range; +use std::sync::Arc; + +use vortex_buffer::BitBuffer; +use vortex_error::VortexExpect; +use vortex_error::VortexResult; +use vortex_error::vortex_bail; +use vortex_error::vortex_ensure; +use vortex_error::vortex_panic; +use vortex_mask::Mask; +use vortex_session::VortexSession; +use vortex_session::registry::CachedId; + +use super::Probe; +use 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::ListViewArray; +use crate::arrays::filter::FilterReduce; +use crate::arrays::filter::FilterReduceAdaptor; +use crate::arrays::slice::SliceReduce; +use crate::arrays::slice::SliceReduceAdaptor; +use crate::buffer::BufferHandle; +use crate::dtype::DType; +use crate::dtype::Nullability; +use crate::match_smallest_list_offset_type; +use crate::optimizer::rules::ParentRuleSet; +use crate::scalar::Scalar; +use crate::scalar_fn::fns::list_contains::ListContainsOptions; +use crate::serde::ArrayChildren; +use crate::validity::Validity; + +/// A [`PreparedSet`]-encoded array. +pub type PreparedSetArray = Array; + +/// A constant, non-null list whose elements are prepared as a set for membership probes. +/// +/// Every row holds the same list, as in a [`ConstantArray`]. When [`ListContains`] gets a constant +/// list and a needle that is not canonical, it puts this array in place of the list and gives the +/// node back to the executor. Thus a [`ListContainsElementKernel`] of the needle encoding gets the +/// prepared set as its list, and can probe its own values with [`PreparedSetData::contains`]. +/// +/// The probe is shared, so a slice or a filter of this array does not build it again. This encoding +/// exists only during execution, and it cannot be serialized. +/// +/// [`ListContains`]: crate::scalar_fn::fns::list_contains::ListContains +/// [`ListContainsElementKernel`]: crate::scalar_fn::fns::list_contains::ListContainsElementKernel +#[derive(Clone, Debug)] +pub struct PreparedSet; + +/// The data of a [`PreparedSetArray`]: the elements of the list, and the probe built from them. +#[derive(Clone)] +pub struct PreparedSetData { + /// All elements of the list, null elements included, with the element dtype of the list. + elements: ArrayRef, + /// The nullability of the list dtype. The list itself is never null. + nullability: Nullability, + pub(super) set: Arc, +} + +/// The non-null elements of the list in a probe structure, and the facts about the list that +/// decide the answer when no element matches. +pub(super) struct ElementSet { + pub(super) probe: Box, + /// Whether the list holds a null element, which a probe does not hold. + has_null_element: bool, + /// Whether the list holds no element at all, counting null elements. + is_empty: bool, +} + +impl PreparedSetData { + /// Prepares `elements`, the elements of a non-null list with the list nullability + /// `nullability`. + pub(super) fn try_new( + elements: ArrayRef, + nullability: Nullability, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let is_empty = elements.is_empty(); + + // A null element never equals a needle, so the probe is better off without it. + let valid = elements.validity()?.execute_mask(elements.len(), ctx)?; + let has_null_element = !valid.all_true(); + let probe_elements = if has_null_element { + elements.filter(valid)? + } else { + elements.clone() + }; + + let probe = new_probe(probe_elements, ctx)?; + + Ok(Self { + elements, + nullability, + set: Arc::new(ElementSet { + probe, + has_null_element, + is_empty, + }), + }) + } + + /// The elements of the list that every row holds. + pub fn elements(&self) -> &ArrayRef { + &self.elements + } + + /// The dtype of the list that every row holds. + fn list_dtype(&self) -> DType { + DType::List(Arc::new(self.elements.dtype().clone()), self.nullability) + } + + /// Whether each of `needles` is an element of the list, under `options`. + /// + /// 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.set.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.set.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.set.is_empty && !options.sql_null_semantics { + Validity::NonNullable + } else if options.sql_null_semantics && self.set.has_null_element { + // Only a match is known. A comparison with the null element makes a non-match unknown. + needle_validity.and(Validity::from(bits.clone()))? + } else { + needle_validity + }; + + Ok(BoolArray::new(bits, validity.union_nullability(nullability)).into_array()) + } + + /// Fails when needles of `needle_dtype` cannot be elements of the list. + fn check_needle_dtype(&self, needle_dtype: &DType) -> VortexResult<()> { + let element_dtype = self.elements.dtype(); + if !element_dtype.eq_ignore_nullability(needle_dtype) { + vortex_bail!( + "Element type {} of list does not match search value {}", + element_dtype, + needle_dtype, + ); + } + Ok(()) + } +} + +impl Debug for PreparedSetData { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.debug_struct("PreparedSetData") + .field("elements", &self.elements) + .field("nullability", &self.nullability) + .finish_non_exhaustive() + } +} + +impl Display for PreparedSetData { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "elements: {}", self.elements.len()) + } +} + +impl ArrayHash for PreparedSetData { + fn array_hash(&self, state: &mut H, accuracy: EqMode) { + self.elements.array_hash(state, accuracy); + self.nullability.hash(state); + } +} + +impl ArrayEq for PreparedSetData { + fn array_eq(&self, other: &Self, accuracy: EqMode) -> bool { + self.nullability == other.nullability && self.elements.array_eq(&other.elements, accuracy) + } +} + +impl Array { + /// Prepares `elements`, the elements of a non-null list with the list nullability + /// `nullability`, 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 the elements cannot be executed into the probe. + pub fn try_new( + elements: ArrayRef, + nullability: Nullability, + len: usize, + ctx: &mut ExecutionCtx, + ) -> VortexResult { + let data = PreparedSetData::try_new(elements, nullability, ctx)?; + Ok(Self::from_data(data, len)) + } + + /// An array of `len` rows that share the prepared set `data`. + fn from_data(data: PreparedSetData, len: usize) -> Self { + let dtype = data.list_dtype(); + + // SAFETY: the dtype is the dtype of the list that every row holds. + unsafe { Array::from_parts_unchecked(ArrayParts::new(PreparedSet, dtype, len, data)) } + } +} + +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. + let data = array.data(); + let n_elements = data.elements.len(); + + // Every row has the same offset and size, so use the narrowest width that fits the list. + let (offsets, sizes) = match_smallest_list_offset_type!(n_elements, |O| { + let size = + O::try_from(n_elements).vortex_expect("list length fits the chosen offset type"); + ( + ConstantArray::new::(O::default(), array.len()).into_array(), + ConstantArray::new::(size, array.len()).into_array(), + ) + }); + + // SAFETY: every view points at the range [0, n_elements) of the elements, and the list + // is never null. + let list = unsafe { + ListViewArray::new_unchecked( + data.elements.clone(), + offsets, + sizes, + Validity::from(data.nullability), + ) + }; + + Ok(ExecutionResult::done(list)) + } + + fn reduce_parent( + 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 { + let elements = (0..array.elements.len()) + .map(|idx| array.elements.execute_scalar(idx, ctx)) + .collect::>>()?; + + Ok(Scalar::list( + array.elements.dtype().clone(), + elements, + array.nullability, + )) + } +} + +impl ValidityVTable for PreparedSet { + fn validity(_array: ArrayView<'_, PreparedSet>) -> VortexResult { + // The list is never null. + Ok(Validity::AllValid) + } +} + +impl SliceReduce for PreparedSet { + fn slice(array: ArrayView<'_, Self>, range: Range) -> VortexResult> { + Ok(Some( + PreparedSetArray::from_data(array.data().clone(), range.len()).into_array(), + )) + } +} + +impl FilterReduce for PreparedSet { + fn filter(array: ArrayView<'_, Self>, mask: &Mask) -> VortexResult> { + Ok(Some( + PreparedSetArray::from_data(array.data().clone(), mask.true_count()).into_array(), + )) + } +} diff --git a/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs new file mode 100644 index 00000000000..0d3028c8e01 --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/mod.rs @@ -0,0 +1,515 @@ +// 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; + +use std::hash::BuildHasher; + +pub use array::PreparedSet; +pub use array::PreparedSetArray; +pub use array::PreparedSetData; +use num_traits::ToPrimitive; +use num_traits::WrappingSub; +use vortex_buffer::BitBuffer; +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()); + let mut short = HashTable::new(); + let mut long = HashTable::new(); + + for (idx, view) in views.iter().enumerate() { + heads.insert(view_head(view)); + + if view.is_inlined() { + let whole = view.as_u128(); + if let HashTableEntry::Vacant(vacant) = short.entry( + hasher.hash_one(whole), + |&other| other == whole, + |&other| hasher.hash_one(other), + ) { + vacant.insert(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..c5ca8c3b174 --- /dev/null +++ b/vortex-array/src/scalar_fn/fns/list_contains/prepared/tests.rs @@ -0,0 +1,458 @@ +// 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::arrays::constant::list_scalar_elements; +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 elements of the non-null list scalar `list`. +fn prepare(list: &Scalar, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + PreparedSetData::try_new(elements, list.dtype().nullability(), ctx) +} + +/// Prepares the elements of the non-null list scalar `list` as a set repeated `len` times. +fn prepare_array(list: &Scalar, len: usize, ctx: &mut ExecutionCtx) -> VortexResult { + let elements = list_scalar_elements(&list.as_list(), ctx.allocator()); + Ok(PreparedSetArray::try_new(elements, list.dtype().nullability(), len, ctx)?.into_array()) +} + +fn nested_needles() -> ArrayRef { + ListArray::try_new( + PrimitiveArray::from_option_iter([Some(1i32), None, Some(2), None, Some(9)]).into_array(), + 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, &mut ctx)?; + let result = set.contains(&needles, &options, &mut ctx)?; + // The old fallback returned a lazy OR tree here. + assert!(result.is::()); + let expected = (0..needles.len()) + .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, &mut ctx)?; + assert!(set.set.probe.is_bitmap()); + let non_match = (!sql_null_semantics).then_some(false); + assert_arrays_eq!( + set.contains(&needles, &options, &mut ctx)?, + 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, &mut ctx)?; + assert_eq!(set.set.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(elements, Nullability::NonNullable, &mut ctx)?; + assert_eq!(set.set.probe.is_bitmap(), dense); + assert_arrays_eq!( + set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, + BoolArray::from_iter(expected), + &mut ctx + ); + Ok(()) +} + +/// The set `{2, null}` of nullable `i32`. +fn set_with_null() -> Scalar { + let element = DType::Primitive(PType::I32, Nullability::Nullable); + 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, &mut ctx)?; + + 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, &mut ctx)?; + 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, &mut ctx)?; + 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, &mut ctx)?; + let needle = ConstantArray::new(member, 4).into_array(); + let result = ListContains::try_new_opts(set, needle, ListContainsOptions::default())? + .into_array() + .execute::(&mut ctx)?; + assert_eq!( + result.as_constant(), + Some(Scalar::bool(true, Nullability::Nullable)) + ); + Ok(()) +} + +#[rstest] +#[case::default(ListContainsOptions::default(), [Some(true), Some(false), None])] +#[case::sql(ListContainsOptions { sql_null_semantics: true }, [Some(true), None, None])] +fn test_result_from_bits_applies_null_semantics( + #[case] options: ListContainsOptions, + #[case] expected: [Option; 3], +) -> VortexResult<()> { + // Against `{2, null}`: a match, a non-match, and a null needle whose bit has no effect. + let mut ctx = array_session().create_execution_ctx(); + let set = prepare(&set_with_null(), &mut ctx)?; + let needle_dtype = DType::Primitive(PType::I32, Nullability::Nullable); + let bits = BitBuffer::from_iter([true, false, true]); + let validity = Validity::from_iter([true, true, false]); + + let result = set.result_from_bits(bits, validity, &needle_dtype, &options)?; + assert_arrays_eq!(result, BoolArray::from_iter(expected), &mut ctx); + Ok(()) +} + +#[test] +fn test_result_from_bits_rejects_mismatched_needles() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let set = prepare(&set_with_null(), &mut ctx)?; + let options = ListContainsOptions::default(); + let bits = BitBuffer::from_iter([true, false]); + + let short_validity = Validity::from_iter([true]); + let needle_dtype = DType::Primitive(PType::I32, Nullability::Nullable); + assert!( + set.result_from_bits(bits.clone(), short_validity, &needle_dtype, &options) + .is_err() + ); + + let wrong_dtype = DType::Primitive(PType::I64, Nullability::Nullable); + assert!( + set.result_from_bits(bits, Validity::AllValid, &wrong_dtype, &options) + .is_err() + ); + Ok(()) +} + +#[test] +fn test_bytes_set_short_and_long_views() -> VortexResult<()> { + let mut ctx = array_session().create_execution_ctx(); + let elements = VarBinViewArray::from_iter_nullable_str([ + Some("ab"), + Some(""), + Some("a long string with suffix A"), + Some("another long string"), + None, + Some("ab"), + ]) + .into_array(); + let set = PreparedSetData::try_new(elements, Nullability::NonNullable, &mut ctx)?; + + // "abc" shares the prefix of a short element, and "... suffix B" shares the head of a long + // one: length and first 4 bytes. Neither is an element. + let needles = VarBinViewArray::from_iter_nullable_str([ + Some("ab"), + Some("abc"), + Some(""), + Some("a long string with suffix B"), + Some("another long string"), + Some("a string never in the set"), + None, + ]) + .into_array(); + + assert_arrays_eq!( + set.contains(&needles, &ListContainsOptions::default(), &mut ctx)?, + BoolArray::from_iter([ + Some(true), + Some(false), + Some(true), + Some(false), + Some(true), + Some(false), + None, + ]), + &mut ctx + ); + Ok(()) +} diff --git a/vortex-jni/src/expression.rs b/vortex-jni/src/expression.rs index 77d709ddd38..154f7056620 100644 --- a/vortex-jni/src/expression.rs +++ b/vortex-jni/src/expression.rs @@ -38,6 +38,7 @@ use vortex::expr::between; use vortex::expr::get_item; use vortex::expr::is_not_null; use vortex::expr::is_null; +use vortex::expr::list_contains_opts; use vortex::expr::lit; use vortex::expr::merge_opts; use vortex::expr::not; @@ -60,6 +61,8 @@ use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; use vortex::scalar_fn::fns::like::Like; use vortex::scalar_fn::fns::like::LikeOptions; +use vortex::scalar_fn::fns::list_contains::ListContainsOptions; +use vortex::scalar_fn::fns::literal::Literal; use vortex::scalar_fn::fns::merge::DuplicateHandling; use vortex::scalar_fn::fns::operators::Operator; @@ -362,6 +365,33 @@ fn strict_from_bool(value: jboolean) -> StrictComparison { } } +/// Build `list_contains(list, needle)`: whether the list-typed `list` expression contains +/// `needle`. +/// +/// `list` must evaluate to a Vortex `List`; the list's element dtype must match `needle`'s dtype +/// ignoring nullability. With a list literal on the left and a column on the right this is a +/// set-membership (`IN`) test that a constant-set kernel answers in one pass over the column. +/// +/// `sql_null_semantics` selects how a null list element behaves: off, it never matches and a +/// non-matching needle is `false`; on, it is SQL's unknown value and a non-matching needle is +/// `null`, so `NOT IN` never admits it. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_listContains( + _env: EnvUnowned, + _class: JClass, + list: jlong, + needle: jlong, + sql_null_semantics: jboolean, +) -> jlong { + let list = unsafe { expr_ref(list) }.clone(); + let needle = unsafe { expr_ref(needle) }.clone(); + into_raw(list_contains_opts( + list, + needle, + ListContainsOptions { sql_null_semantics }, + )) +} + #[unsafe(no_mangle)] pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalBool( _env: EnvUnowned, @@ -437,6 +467,93 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalBinary( }) } +/// Build a non-empty list literal out of the literal expressions in `elements`. +/// +/// Every element must be a literal (`vortex.literal`) whose dtype matches the first element's +/// ignoring nullability; the list's element dtype is that shared dtype, made nullable if any +/// element is nullable, and each element is cast to it. The list itself is non-nullable — use +/// [`Java_dev_vortex_jni_NativeExpression_literalEmptyList`] for a null or empty list, which +/// cannot infer an element dtype from its (absent) elements. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalList( + mut env: EnvUnowned, + _class: JClass, + elements: JLongArray, +) -> jlong { + try_or_throw(&mut env, |env| { + let ptrs = unsafe { elements.get_elements(env, ReleaseMode::NoCopyBack) }?; + let scalars = ptrs + .iter() + .map(|ptr| literal_scalar(unsafe { expr_ref(*ptr) })) + .collect::, _>>()?; + Ok(into_raw(lit(list_scalar(&scalars)?))) + }) +} + +/// The scalar behind a literal expression, or an error if the expression is not a literal. +fn literal_scalar(expr: &Expression) -> Result { + expr.as_opt::().cloned().ok_or_else(|| { + vortex_err!("list literal elements must themselves be literals, got {expr}").into() + }) +} + +/// Collect literal scalars into a single non-nullable list scalar. +fn list_scalar(elements: &[Scalar]) -> Result { + let Some(first) = elements.first() else { + throw_runtime!("list literal requires at least one element; use an empty list literal"); + }; + + let mut nullability = Nullability::NonNullable; + for element in elements { + if !element.dtype().eq_ignore_nullability(first.dtype()) { + throw_runtime!( + "list literal elements must share a dtype, got {} and {}", + first.dtype(), + element.dtype() + ); + } + nullability |= element.dtype().nullability(); + } + + let element_dtype = first.dtype().with_nullability(nullability); + let children = elements + .iter() + .map(|element| element.cast(&element_dtype)) + .collect::, _>>()?; + Ok(Scalar::list( + Arc::new(element_dtype), + children, + Nullability::NonNullable, + )) +} + +/// Build an empty (or null) list literal whose element dtype is selected by `element_dtype_tag`. +/// +/// The tag table is the one [`Java_dev_vortex_jni_NativeExpression_literalNull`] reads; see +/// `dev.vortex.api.Expression.DType` on the Java side for the source of truth. Elements are +/// nullable so that the literal accepts a nullable column as its needle. +#[unsafe(no_mangle)] +pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalEmptyList( + mut env: EnvUnowned, + _class: JClass, + element_dtype_tag: jbyte, + is_null_flag: jboolean, +) -> jlong { + try_or_throw(&mut env, |_| { + let element_dtype = Arc::new(parse_null_dtype(element_dtype_tag)?); + if is_null_flag { + return Ok(into_raw(lit(Scalar::null(DType::List( + element_dtype, + Nullability::Nullable, + ))))); + } + Ok(into_raw(lit(Scalar::list_empty( + element_dtype, + Nullability::NonNullable, + )))) + }) +} + /// Build a decimal literal from a two's-complement big-endian byte representation of the /// unscaled value (the format produced by Java's `BigInteger.toByteArray()`). #[unsafe(no_mangle)] @@ -666,10 +783,26 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalUuid( }) } -/// Build a typed null literal whose nullable dtype is selected by `dtype_tag`. +/// Parse a nullable primitive [`DType`] from the wire-encoded byte tag. /// /// Tag values intentionally do not overlap with [`parse_time_unit`]. /// See `dev.vortex.api.Expression.DType` on the Java side for the source of truth. +fn parse_null_dtype(tag: jbyte) -> Result { + Ok(match tag { + 0 => DType::Bool(Nullability::Nullable), + 1 => DType::Primitive(PType::I8, Nullability::Nullable), + 2 => DType::Primitive(PType::I16, Nullability::Nullable), + 3 => DType::Primitive(PType::I32, Nullability::Nullable), + 4 => DType::Primitive(PType::I64, Nullability::Nullable), + 5 => DType::Primitive(PType::F32, Nullability::Nullable), + 6 => DType::Primitive(PType::F64, Nullability::Nullable), + 7 => DType::Utf8(Nullability::Nullable), + 8 => DType::Binary(Nullability::Nullable), + other => throw_runtime!("unknown null dtype tag: {other}"), + }) +} + +/// Build a typed null literal whose nullable dtype is selected by `dtype_tag`. #[unsafe(no_mangle)] pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalNull( mut env: EnvUnowned, @@ -677,18 +810,6 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalNull( dtype_tag: jbyte, ) -> jlong { try_or_throw(&mut env, |_| { - let dtype = match dtype_tag { - 0 => DType::Bool(Nullability::Nullable), - 1 => DType::Primitive(PType::I8, Nullability::Nullable), - 2 => DType::Primitive(PType::I16, Nullability::Nullable), - 3 => DType::Primitive(PType::I32, Nullability::Nullable), - 4 => DType::Primitive(PType::I64, Nullability::Nullable), - 5 => DType::Primitive(PType::F32, Nullability::Nullable), - 6 => DType::Primitive(PType::F64, Nullability::Nullable), - 7 => DType::Utf8(Nullability::Nullable), - 8 => DType::Binary(Nullability::Nullable), - other => throw_runtime!("unknown null dtype tag: {other}"), - }; - Ok(into_raw(lit(Scalar::null(dtype)))) + Ok(into_raw(lit(Scalar::null(parse_null_dtype(dtype_tag)?)))) }) } diff --git a/vortex-python/python/vortex/_lib/expr.pyi b/vortex-python/python/vortex/_lib/expr.pyi index 6cd2b58bfd5..7544b6f830d 100644 --- a/vortex-python/python/vortex/_lib/expr.pyi +++ b/vortex-python/python/vortex/_lib/expr.pyi @@ -3,14 +3,14 @@ from collections.abc import Iterable, Mapping, Sequence from datetime import date, datetime -from typing import Literal, TypeAlias, final +from typing import Any, Literal, TypeAlias, final from typing_extensions import override from .dtype import DType -from .scalar import ScalarPyType +from .scalar import Scalar, ScalarPyType -IntoExpr: TypeAlias = Expr | bool | int | float | str | bytes | date | datetime | None +IntoExpr: TypeAlias = Expr | Scalar | ScalarPyType | list[Any] | dict[str, Any] | date | datetime """A value accepted anywhere an expression is expected. Non-``Expr`` values become literals.""" VariantPath: TypeAlias = str | int | Sequence[str | int] @@ -46,7 +46,7 @@ class Expr: # Leaves and scope def root() -> Expr: ... def column(name: str) -> Expr: ... -def literal(dtype: DType, value: ScalarPyType) -> Expr: ... +def literal(dtype: DType, value: object) -> Expr: ... def get_item(field: str, child: IntoExpr | None = None) -> Expr: ... # Boolean logic @@ -103,7 +103,8 @@ def merge( ) -> Expr: ... # Lists -def list_contains(child: IntoExpr, value: IntoExpr) -> Expr: ... +def list_contains(child: IntoExpr, value: IntoExpr, *, sql_null_semantics: bool = False) -> Expr: ... +def in_list(value: IntoExpr, list: IntoExpr) -> Expr: ... def list_length(child: IntoExpr) -> Expr: ... def list_sum(child: IntoExpr, *, skip_nans: bool = True) -> Expr: ... diff --git a/vortex-python/python/vortex/expr.py b/vortex-python/python/vortex/expr.py index d02ed5ebab6..afd78cb4d85 100644 --- a/vortex-python/python/vortex/expr.py +++ b/vortex-python/python/vortex/expr.py @@ -21,6 +21,7 @@ gt, gt_eq, ilike, + in_list, is_not_null, is_null, like, @@ -67,6 +68,7 @@ "gt", "gt_eq", "ilike", + "in_list", "is_not_null", "is_null", "like", diff --git a/vortex-python/src/expr/mod.rs b/vortex-python/src/expr/mod.rs index 8d76d20b051..806aa5ddcd5 100644 --- a/vortex-python/src/expr/mod.rs +++ b/vortex-python/src/expr/mod.rs @@ -10,7 +10,6 @@ use pyo3::intern; use pyo3::prelude::*; use pyo3::types::*; use vortex::aggregate_fn::NumericalAggregateOpts; -use vortex::dtype::DType; use vortex::dtype::FieldName; use vortex::dtype::FieldNames; use vortex::dtype::Nullability; @@ -23,6 +22,7 @@ use vortex::scalar_fn::ScalarFnVTableExt; use vortex::scalar_fn::fns::between::BetweenOptions; use vortex::scalar_fn::fns::between::StrictComparison; use vortex::scalar_fn::fns::binary::Binary; +use vortex::scalar_fn::fns::list_contains::ListContainsOptions; use vortex::scalar_fn::fns::merge::DuplicateHandling; use vortex::scalar_fn::fns::operators::Operator; use vortex::scalar_fn::fns::variant_get::VariantPath; @@ -85,6 +85,7 @@ pub(crate) fn init(py: Python, parent: &Bound) -> PyResult<()> { // Lists m.add_function(wrap_pyfunction!(list_contains, &m)?)?; + m.add_function(wrap_pyfunction!(in_list, &m)?)?; m.add_function(wrap_pyfunction!(list_length, &m)?)?; m.add_function(wrap_pyfunction!(list_sum, &m)?)?; @@ -502,7 +503,8 @@ pub fn literal<'py>( dtype: &Bound<'py, PyDType>, value: &Bound<'py, PyAny>, ) -> PyResult> { - scalar(dtype.borrow().inner().clone(), value) + let scalar = scalar_helper(value, Some(dtype.borrow().inner())).map_err(PyErr::from)?; + Bound::new(value.py(), PyExpr { inner: lit(scalar) }) } /// Create an expression that refers to the identity scope. @@ -593,16 +595,6 @@ pub fn get_item(field: String, child: Option) -> PyExpr { } } -pub fn scalar<'py>(dtype: DType, value: &Bound<'py, PyAny>) -> PyResult> { - let py = value.py(); - Bound::new( - py, - PyExpr { - inner: lit(scalar_helper(value, Some(&dtype))?), - }, - ) -} - /// Negate a Boolean expression. /// /// Parameters @@ -1097,20 +1089,64 @@ pub fn merge(exprs: &Bound<'_, PyAny>, duplicate_handling: &str) -> PyResult PyExpr { +#[pyo3(signature = (child, value, *, sql_null_semantics = false))] +pub fn list_contains(child: PyIntoExpr, value: PyIntoExpr, sql_null_semantics: bool) -> PyExpr { + PyExpr { + inner: expr::list_contains_opts( + child.into_inner(), + value.into_inner(), + ListContainsOptions { sql_null_semantics }, + ), + } +} + +/// SQL ``value IN (list)`` with SQL null semantics. +/// +/// A null value produces null. A non-match also produces null if the list contains null. +/// Use ``~in_list(value, list)`` for SQL ``NOT IN``. +/// +/// Parameters +/// ---------- +/// value : :class:`Any` +/// The value to search for. +/// list : :class:`Any` +/// A list expression or a Python list. Use :func:`.literal` with an explicit list dtype for +/// empty or null lists and for element types other than the inferred Python scalar types. +/// +/// Returns +/// ------- +/// :class:`vortex.Expr` +/// +/// Examples +/// -------- +/// +/// ```python +/// >>> import vortex.expr as ve +/// >>> ve.in_list(ve.column("age"), [25, 30]) +/// +/// ``` +#[pyfunction] +#[pyo3(signature = (value, list))] +pub fn in_list(value: PyIntoExpr, list: PyIntoExpr) -> PyExpr { PyExpr { - inner: expr::list_contains(child.into_inner(), value.into_inner()), + inner: expr::in_list(value.into_inner(), list.into_inner()), } } diff --git a/vortex-python/src/scalar/factory.rs b/vortex-python/src/scalar/factory.rs index 13e46b048c9..968e45e55c9 100644 --- a/vortex-python/src/scalar/factory.rs +++ b/vortex-python/src/scalar/factory.rs @@ -22,6 +22,7 @@ use vortex::scalar::DecimalValue; use vortex::scalar::Scalar; use crate::dtype::PyDType; +use crate::error::PyVortexError; use crate::error::PyVortexResult; use crate::scalar::PyScalar; use crate::scalar::bool; @@ -180,32 +181,44 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes if let Some(DType::List(element_dtype, ..)) = dtype { let elements = list .iter() - .map(|e| scalar_helper_inner(&e, Some(element_dtype))) - .try_collect()?; - Scalar::list( + .map(|e| scalar_helper(&e, Some(element_dtype))) + .collect::>>()?; + return Ok(Scalar::list( Arc::clone(element_dtype), elements, Nullability::NonNullable, - ); + )); } else { - // If no dtype was provided, we need to infer the element dtype from the list contents. - // We do this in a greedy way taking the first element dtype we find. - let mut elements = Vec::with_capacity(list.len()); - let mut element_dtype = None; - - for element in list.iter() { - let scalar = scalar_helper_inner(&element, element_dtype.as_ref())?; - if element_dtype.is_none() { - element_dtype = Some(scalar.dtype().clone()); + let elements = list + .iter() + .map(|element| scalar_helper_inner(&element, None)) + .collect::>>()?; + let element_dtype = elements + .iter() + .find(|element| !matches!(element.dtype(), DType::Null)) + .map(|element| element.dtype().clone()) + .unwrap_or(DType::Null); + let mut nullability = element_dtype.nullability(); + for element in &elements { + if !matches!(element.dtype(), DType::Null) + && !element.dtype().eq_ignore_nullability(&element_dtype) + { + return Err(PyValueError::new_err(format!( + "list elements must share a dtype, got {} and {}", + element_dtype, + element.dtype() + ))); } - elements.push(scalar); + nullability |= element.dtype().nullability(); } + let element_dtype = element_dtype.with_nullability(nullability); + let elements = elements + .iter() + .map(|element| element.cast(&element_dtype).map_err(PyVortexError::from)) + .collect::>>()?; return Ok(Scalar::list( - element_dtype - .map(Arc::new) - // Empty list defaults to Null dtype - .unwrap_or_else(|| Arc::new(DType::Null)), + Arc::new(element_dtype), elements, Nullability::NonNullable, )); diff --git a/vortex-python/test/test_expr.py b/vortex-python/test/test_expr.py index ed5cf3bee76..959de013b55 100644 --- a/vortex-python/test/test_expr.py +++ b/vortex-python/test/test_expr.py @@ -84,6 +84,8 @@ def column_values(vxf: vx.VortexFile, projection: Expr) -> list[object]: "merge": lambda: ve.merge([ve.select(["name"]), ve.select(["age"])]), "merge_rightmost": lambda: ve.merge([ve.select(["name"]), ve.select(["name"])], duplicate_handling="rightmost"), "list_contains": lambda: ve.list_contains(ve.column("scores"), 5), + "list_contains_sql": lambda: ve.list_contains(ve.column("scores"), 5, sql_null_semantics=True), + "in_list": lambda: ve.in_list(ve.column("age"), [25, 30]), "list_length": lambda: ve.list_length(ve.column("scores")), "list_sum": lambda: ve.list_sum(ve.column("scores")), "list_sum_nans": lambda: ve.list_sum(ve.column("scores"), skip_nans=False), @@ -209,6 +211,46 @@ def test_arithmetic_and_list_functions(people: vx.VortexFile) -> None: assert column_values(people, ve.get_item("city", ve.column("nested"))) == ["Paris", "Berlin", "Paris", "Lima"] +@pytest.mark.parametrize( + "values,default,sql", + [ + ([30, 57], [True, False, None, True], [True, False, None, True]), + ([30, None], [True, False, None, False], [True, None, None, None]), + ([None, 30], [True, False, None, False], [True, None, None, None]), + ([None], [False, False, None, False], [None, None, None, None]), + ([], [False, False, False, False], [False, False, None, False]), + (None, [None, None, None, None], [None, None, None, None]), + ], +) +def test_list_membership_null_semantics( + people: vx.VortexFile, + values: list[int | None] | None, + default: list[bool | None], + sql: list[bool | None], +): + members = ve.literal(vx.list_(vx.int_(64, nullable=True), nullable=True), values) + age = ve.column("age") + assert column_values(people, ve.list_contains(members, age)) == default + assert column_values(people, ve.list_contains(members, age, sql_null_semantics=True)) == sql + expr = ve.in_list(age, members) + assert column_values(people, expr) == sql + assert column_values(people, ve.deserialize(expr.serialize())) == sql + assert column_values(people, ~expr) == [None if value is None else not value for value in sql] + + +@pytest.mark.parametrize("members", [[30, None], [None, 30]]) +def test_in_list_python_list_filter(people: vx.VortexFile, members: list[int | None]): + expr = ve.in_list(ve.column("age"), members) + assert names(people, expr) == ["Alice"] + assert names(people, ~expr) == [] + assert names(people, ~ve.in_list(ve.column("age"), [30])) == ["Bob", "Charlie"] + + +def test_in_list_large_set_and_strings(people: vx.VortexFile): + assert names(people, ve.in_list(ve.column("age"), list(range(1000)))) == ["Alice", "Bob", "Charlie"] + assert names(people, ve.in_list(ve.column("name"), ["Alice", "Charlie"])) == ["Alice", "Charlie"] + + def test_case_when_semantics(people: vx.VortexFile) -> None: expr = ve.case_when([(ve.gt(ve.column("age"), 40), "senior"), (ve.gt(ve.column("age"), 26), "mid")], "junior") assert column_values(people, expr) == ["mid", "junior", "junior", "senior"] diff --git a/vortex-python/test/test_scalar.py b/vortex-python/test/test_scalar.py index 24c5889c9a4..9af210dc34b 100644 --- a/vortex-python/test/test_scalar.py +++ b/vortex-python/test/test_scalar.py @@ -39,6 +39,20 @@ def test_f16() -> None: assert scalar.as_py() == 1.0 +@pytest.mark.parametrize("values", [[1, None], [None, 1], [None, None], []]) +def test_list_scalar_nullability(values: list[int | None]): + assert vx.scalar(values).as_py() == values + dtype = vx.list_(vx.int_(32, nullable=True)) + scalar = vx.scalar(values, dtype=dtype) + assert scalar.dtype == dtype + assert scalar.as_py() == values + + +def test_list_scalar_rejects_mixed_types(): + with pytest.raises(ValueError, match="must share a dtype"): + _ = vx.scalar([1, "two"]) + + @pytest.mark.parametrize( "precision,scale,stored,expected", [