Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions docs/api/python/expr.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
73 changes: 73 additions & 0 deletions java/vortex-jni/src/main/java/dev/vortex/api/Expression.java
Original file line number Diff line number Diff line change
Expand Up @@ -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}.
*
* <p>{@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));
}
Expand Down Expand Up @@ -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.
*
* <p>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()}).
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand All @@ -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);
Expand Down
51 changes: 51 additions & 0 deletions java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down
Original file line number Diff line number Diff line change
@@ -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.
*
* <p>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.
*
* <p>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.
*
* <p>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<Integer> 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<Integer> 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;
}
}
2 changes: 1 addition & 1 deletion vortex-array/benches/list_contains_set.rs
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,7 @@ fn random_i64(len: usize) -> (Vec<i64>, Vec<i64>) {

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()
Expand Down
Loading
Loading