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",
[