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..1ae991b9839 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
@@ -119,6 +119,20 @@ public static Expression binary(BinaryOp op, Expression lhs, Expression rhs) {
return new Expression(NativeExpression.binary(op.code(), lhs.nativePointer(), rhs.nativePointer()));
}
+ /**
+ * Apply a spatial function to native geometry operands, each a column or a geometry literal (see
+ * {@link #literalGeometry(byte[])}). The number of operands must match the function's arity.
+ */
+ public static Expression spatial(SpatialFunction function, Expression... operands) {
+ Preconditions.checkArgument(
+ operands.length == function.arity(),
+ "%s takes %s operands, got %s",
+ function,
+ function.arity(),
+ operands.length);
+ return new Expression(NativeExpression.spatial(function.code(), nativePointers(operands)));
+ }
+
public static Expression not(Expression child) {
return new Expression(NativeExpression.not(child.nativePointer()));
}
@@ -178,6 +192,45 @@ public static Expression literal(long value) {
return new Expression(NativeExpression.literalI64(value, false));
}
+ /** Create an unsigned 8-bit integer literal. {@code value} must be in {@code [0, 255]}. */
+ public static Expression literalU8(int value) {
+ Preconditions.checkArgument(value >= 0 && value <= 0xFF, "u8 literal out of range: %s", value);
+ return new Expression(NativeExpression.literalU8((byte) value, false));
+ }
+
+ /** Create an unsigned 16-bit integer literal. {@code value} must be in {@code [0, 65535]}. */
+ public static Expression literalU16(int value) {
+ Preconditions.checkArgument(value >= 0 && value <= 0xFFFF, "u16 literal out of range: %s", value);
+ return new Expression(NativeExpression.literalU16((short) value, false));
+ }
+
+ /** Create an unsigned 32-bit integer literal. {@code value} must be in {@code [0, 2^32 - 1]}. */
+ public static Expression literalU32(long value) {
+ Preconditions.checkArgument(value >= 0 && value <= 0xFFFF_FFFFL, "u32 literal out of range: %s", value);
+ return new Expression(NativeExpression.literalU32((int) value, false));
+ }
+
+ /**
+ * Create an unsigned 64-bit integer literal from its bit pattern, so values above {@link Long#MAX_VALUE} are passed
+ * as negative longs (for example via {@link Long#parseUnsignedLong(String)}).
+ */
+ public static Expression literalU64(long bits) {
+ return new Expression(NativeExpression.literalU64(bits, false));
+ }
+
+ /** Create an unsigned 64-bit integer literal. {@code value} must be in {@code [0, 2^64 - 1]}. */
+ public static Expression literalU64(BigInteger value) {
+ Preconditions.checkArgument(value != null, "use nullLiteral(DType.U64) for a null u64 literal");
+ Preconditions.checkArgument(
+ value.signum() >= 0 && value.bitLength() <= Long.SIZE, "u64 literal out of range: %s", value);
+ return literalU64(value.longValue());
+ }
+
+ /** Create a half-precision float literal, rounding {@code value} to the nearest representable half. */
+ public static Expression literalF16(float value) {
+ return new Expression(NativeExpression.literalF16(value, false));
+ }
+
public static Expression literal(float value) {
return new Expression(NativeExpression.literalF32(value, false));
}
@@ -238,6 +291,34 @@ public static Expression nullLiteralTimestamp(TimeUnit unit, String timezone) {
return new Expression(NativeExpression.literalTimestamp(0L, unit.tag(), timezone, true));
}
+ /**
+ * Create a Time (time-of-day) literal. The {@code value} is the number of {@code unit} units since midnight.
+ *
+ * @param unit any unit except {@link TimeUnit#DAYS}. {@link TimeUnit#SECONDS} and {@link TimeUnit#MILLISECONDS}
+ * values must fit in an {@code int}.
+ */
+ public static Expression literalTime(long value, TimeUnit unit) {
+ return new Expression(NativeExpression.literalTime(value, unit.tag(), false));
+ }
+
+ /** Null Time literal. See {@link #literalTime(long, TimeUnit)} for the {@code unit} constraints. */
+ public static Expression nullLiteralTime(TimeUnit unit) {
+ return new Expression(NativeExpression.literalTime(0L, unit.tag(), true));
+ }
+
+ /**
+ * Create a geometry literal from its OGC Well-Known Binary (WKB) encoding, for use with Vortex's spatial functions
+ * and predicate pushdown over geometry columns.
+ *
+ *
The value is decoded into the native Vortex geometry type matching its kind: Point, LineString, Polygon,
+ * MultiPoint, MultiLineString or MultiPolygon, in XY with no coordinate reference system. Geometry collections and
+ * malformed WKB are rejected.
+ */
+ public static Expression literalGeometry(byte[] wkb) {
+ Preconditions.checkArgument(wkb != null, "geometry literal WKB must not be null");
+ return new Expression(NativeExpression.literalGeometry(wkb));
+ }
+
/**
* Create a UUID literal, enabling predicate pushdown over UUID columns. The value is stored as its 16-byte
* big-endian (network order) representation, matching Vortex's UUID extension type and Arrow's canonical UUID type.
@@ -281,6 +362,89 @@ public static Expression nullLiteral(DType dtype) {
return new Expression(NativeExpression.literalNull(dtype.tag()));
}
+ /**
+ * Create a list literal.
+ *
+ *
Composite literals are assembled from other literal expressions. The element dtype is the dtype of the
+ * {@code elementType} literal (typically a typed null such as {@code nullLiteral(DType.I32)}) with nullability
+ * {@code elementsNullable}, and every element, which must itself be a literal, is cast to it.
+ */
+ public static Expression literalList(Expression elementType, boolean elementsNullable, Expression... elements) {
+ return new Expression(NativeExpression.literalList(
+ nativePointers(elements), elementType.nativePointer(), elementsNullable, false));
+ }
+
+ /** Null list literal. See {@link #literalList(Expression, boolean, Expression...)} for the element dtype. */
+ public static Expression nullLiteralList(Expression elementType, boolean elementsNullable) {
+ return new Expression(
+ NativeExpression.literalList(new long[0], elementType.nativePointer(), elementsNullable, true));
+ }
+
+ /**
+ * Create a fixed-size list literal whose size is the number of {@code elements}. The element dtype follows
+ * {@link #literalList(Expression, boolean, Expression...)}.
+ */
+ public static Expression literalFixedSizeList(
+ Expression elementType, boolean elementsNullable, Expression... elements) {
+ return new Expression(NativeExpression.literalFixedSizeList(
+ nativePointers(elements), elementType.nativePointer(), elementsNullable, elements.length, false));
+ }
+
+ /** Null fixed-size list literal of {@code size} elements. */
+ public static Expression nullLiteralFixedSizeList(Expression elementType, boolean elementsNullable, int size) {
+ Preconditions.checkArgument(size >= 0, "fixed-size list size must not be negative: %s", size);
+ return new Expression(NativeExpression.literalFixedSizeList(
+ new long[0], elementType.nativePointer(), elementsNullable, size, true));
+ }
+
+ /** Create a struct literal. Each field must be a literal expression; its dtype becomes the field's dtype. */
+ public static Expression literalStruct(String[] fieldNames, Expression[] fields) {
+ Preconditions.checkArgument(
+ fieldNames.length == fields.length,
+ "struct literal has %s field names but %s fields",
+ fieldNames.length,
+ fields.length);
+ return new Expression(NativeExpression.literalStruct(fieldNames, nativePointers(fields), false));
+ }
+
+ /**
+ * Create a null struct literal. {@code fieldTypes} are literals (typically typed nulls) whose dtypes become the
+ * struct's field dtypes.
+ */
+ public static Expression nullLiteralStruct(String[] fieldNames, Expression[] fieldTypes) {
+ Preconditions.checkArgument(
+ fieldNames.length == fieldTypes.length,
+ "struct literal has %s field names but %s field types",
+ fieldNames.length,
+ fieldTypes.length);
+ return new Expression(NativeExpression.literalStruct(fieldNames, nativePointers(fieldTypes), true));
+ }
+
+ /**
+ * Create a map literal from parallel arrays of literal keys and values.
+ *
+ *
Keys are cast to the non-nullable dtype of the {@code keyType} literal, and values to the dtype of the
+ * {@code valueType} literal with nullability {@code valuesNullable}. Keys are not asserted to be sorted.
+ */
+ public static Expression literalMap(
+ Expression keyType, Expression valueType, boolean valuesNullable, Expression[] keys, Expression[] values) {
+ Preconditions.checkArgument(
+ keys.length == values.length, "map literal has %s keys but %s values", keys.length, values.length);
+ return new Expression(NativeExpression.literalMap(
+ nativePointers(keys),
+ nativePointers(values),
+ keyType.nativePointer(),
+ valueType.nativePointer(),
+ valuesNullable,
+ false));
+ }
+
+ /** Null map literal. See {@link #literalMap} for how the key and value dtypes are derived. */
+ public static Expression nullLiteralMap(Expression keyType, Expression valueType, boolean valuesNullable) {
+ return new Expression(NativeExpression.literalMap(
+ new long[0], new long[0], keyType.nativePointer(), valueType.nativePointer(), valuesNullable, true));
+ }
+
private static long[] nativePointers(Expression[] exprs) {
return Arrays.stream(exprs).mapToLong(Expression::nativePointer).toArray();
}
@@ -311,6 +475,44 @@ public byte code() {
}
}
+ /** Spatial functions over native geometries; codes must match the Rust {@code spatial} table. */
+ public enum SpatialFunction {
+ /** Planar area of each geometry; zero for points and line strings. */
+ AREA((byte) 0, 1),
+ /** Collect each list of homogeneous geometries into the matching multi-geometry. */
+ COLLECT((byte) 1, 1),
+ /** Whether the second geometry lies completely inside the first. */
+ CONTAINS((byte) 2, 2),
+ /** Convex hull of each multipoint, as a polygon. */
+ CONVEX_HULL((byte) 3, 1),
+ /** Planar (Euclidean) distance between two geometries. */
+ DISTANCE((byte) 4, 2),
+ /** Axis-aligned bounding box of each geometry. */
+ ENVELOPE((byte) 5, 1),
+ /** Whether two geometries intersect; boundary contact counts. */
+ INTERSECTS((byte) 6, 2),
+ /** Planar length of each lineal geometry. */
+ LENGTH((byte) 7, 1),
+ /** Line string from two points. */
+ MAKE_LINE((byte) 8, 2);
+
+ private final byte code;
+ private final int arity;
+
+ SpatialFunction(byte code, int arity) {
+ this.code = code;
+ this.arity = arity;
+ }
+
+ public byte code() {
+ return code;
+ }
+
+ public int arity() {
+ return arity;
+ }
+ }
+
/**
* Strategy for resolving duplicate field names in {@link #merge(DuplicateHandling, Expression...)}. Tag values must
* match the Rust {@code parse_duplicate_handling} table.
@@ -332,7 +534,7 @@ public byte tag() {
}
}
- /** Time units for Date/Timestamp literals. Tag values must match the Rust {@code parse_time_unit} table. */
+ /** Time units for Date/Time/Timestamp literals. Tag values must match the Rust {@code parse_time_unit} table. */
public enum TimeUnit {
NANOSECONDS((byte) 0),
MICROSECONDS((byte) 1),
@@ -361,7 +563,12 @@ public enum DType {
F32((byte) 5),
F64((byte) 6),
UTF8((byte) 7),
- BINARY((byte) 8);
+ BINARY((byte) 8),
+ U8((byte) 9),
+ U16((byte) 10),
+ U32((byte) 11),
+ U64((byte) 12),
+ F16((byte) 13);
private final byte tag;
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..cd1d9972a9c 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
@@ -29,6 +29,8 @@ private NativeExpression() {}
public static native long binary(byte operator, long lhs, long rhs);
+ public static native long spatial(byte function, long[] operands);
+
public static native long not(long childPointer);
public static native long isNull(long childPointer);
@@ -50,6 +52,16 @@ public static native long between(
public static native long literalI64(long value, boolean isNull);
+ public static native long literalU8(byte bits, boolean isNull);
+
+ public static native long literalU16(short bits, boolean isNull);
+
+ public static native long literalU32(int bits, boolean isNull);
+
+ public static native long literalU64(long bits, boolean isNull);
+
+ public static native long literalF16(float value, boolean isNull);
+
public static native long literalF32(float value, boolean isNull);
public static native long literalF64(double value, boolean isNull);
@@ -64,9 +76,29 @@ public static native long between(
public static native long literalTimestamp(long value, byte timeUnitTag, String timezone, boolean isNull);
+ public static native long literalTime(long value, byte timeUnitTag, boolean isNull);
+
public static native long literalUuid(byte[] bigEndianBytes, boolean isNull);
+ public static native long literalGeometry(byte[] wkb);
+
public static native long literalNull(byte dtypeTag);
+ public static native long literalList(
+ long[] elementPointers, long elementTypePointer, boolean elementsNullable, boolean isNull);
+
+ public static native long literalFixedSizeList(
+ long[] elementPointers, long elementTypePointer, boolean elementsNullable, int size, boolean isNull);
+
+ public static native long literalStruct(String[] fieldNames, long[] fieldPointers, boolean isNull);
+
+ public static native long literalMap(
+ long[] keyPointers,
+ long[] valuePointers,
+ long keyTypePointer,
+ long valueTypePointer,
+ boolean valuesNullable,
+ boolean isNull);
+
public static native void free(long pointer);
}
diff --git a/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTagParityTest.java b/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTagParityTest.java
index 00712632434..deddff1ac1b 100644
--- a/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTagParityTest.java
+++ b/java/vortex-jni/src/test/java/dev/vortex/api/ExpressionTagParityTest.java
@@ -11,6 +11,7 @@
import dev.vortex.api.Expression.BinaryOp;
import dev.vortex.api.Expression.DType;
import dev.vortex.api.Expression.DuplicateHandling;
+import dev.vortex.api.Expression.SpatialFunction;
import dev.vortex.api.Expression.TimeUnit;
import dev.vortex.jni.NativeExpression;
import dev.vortex.jni.NativeLoader;
@@ -20,13 +21,13 @@
/**
* Parity tests for the byte tags {@link Expression} hands to the native side.
*
- *
Four enums carry a tag that the Rust side switches on: {@link BinaryOp} against {@code parse_op},
+ *
Five enums carry a tag that the Rust side switches on: {@link BinaryOp} against {@code parse_op},
* {@link DuplicateHandling} against {@code parse_duplicate_handling}, {@link TimeUnit} against
- * {@code TimeUnit::try_from}, and {@link DType} against the table in {@code literalNull}. Three of the four say in
- * their javadoc that the values must match the Rust table, but nothing checked it. Drift compiles on both sides, and
- * the two failure modes are not equally loud: a tag past the end of a table reaches the {@code other =>} arm and
- * throws, while a tag that collides with a sibling decodes to the wrong operator or the wrong time unit and returns an
- * expression that reads valid.
+ * {@code TimeUnit::try_from}, {@link DType} against the table in {@code literalNull}, and {@link SpatialFunction}
+ * against the table in {@code spatial}. Four of the five say in their javadoc that the values must match the Rust
+ * table, but nothing checked it. Drift compiles on both sides, and the two failure modes are not equally loud: a tag
+ * past the end of a table reaches the {@code other =>} arm and throws, while a tag that collides with a sibling decodes
+ * to the wrong operator or the wrong time unit and returns an expression that reads valid.
*
*
So each table is pinned twice: the constants against the bytes Rust matches, and every constant against the native
* call that consumes it. Both temporal types are exercised because between them they reject enough units to tell the
@@ -80,6 +81,43 @@ public void aBinaryOpCodePastTheTableIsRejectedByName() {
() -> "unexpected message: " + exception.getMessage());
}
+ @Test
+ public void spatialFunctionCodesMatchTheRustTable() {
+ assertEquals(0, SpatialFunction.AREA.code());
+ assertEquals(1, SpatialFunction.COLLECT.code());
+ assertEquals(2, SpatialFunction.CONTAINS.code());
+ assertEquals(3, SpatialFunction.CONVEX_HULL.code());
+ assertEquals(4, SpatialFunction.DISTANCE.code());
+ assertEquals(5, SpatialFunction.ENVELOPE.code());
+ assertEquals(6, SpatialFunction.INTERSECTS.code());
+ assertEquals(7, SpatialFunction.LENGTH.code());
+ assertEquals(8, SpatialFunction.MAKE_LINE.code());
+ assertEquals(9, SpatialFunction.values().length);
+ }
+
+ @Test
+ public void everySpatialFunctionIsAcceptedWithItsArity() {
+ // The native side rejects an operand count that does not match the function's arity, so a Java arity
+ // that drifts from the Rust signature fails here.
+ for (SpatialFunction function : SpatialFunction.values()) {
+ Expression[] operands = new Expression[function.arity()];
+ for (int i = 0; i < operands.length; i++) {
+ operands[i] = Expression.column("g" + i);
+ }
+ assertNotNull(Expression.spatial(function, operands), () -> "native side rejected " + function);
+ }
+ }
+
+ @Test
+ public void aSpatialFunctionCodePastTheTableIsRejectedByName() {
+ RuntimeException exception = assertThrows(
+ RuntimeException.class,
+ () -> NativeExpression.spatial((byte) SpatialFunction.values().length, new long[0]));
+ assertTrue(
+ exception.getMessage().contains("unknown spatial function code: 9"),
+ () -> "unexpected message: " + exception.getMessage());
+ }
+
@Test
public void duplicateHandlingTagsMatchTheRustTable() {
assertEquals(0, DuplicateHandling.RIGHT_MOST.tag());
@@ -158,7 +196,12 @@ public void nullLiteralDTypeTagsMatchTheRustTable() {
assertEquals(6, DType.F64.tag());
assertEquals(7, DType.UTF8.tag());
assertEquals(8, DType.BINARY.tag());
- assertEquals(9, DType.values().length);
+ assertEquals(9, DType.U8.tag());
+ assertEquals(10, DType.U16.tag());
+ assertEquals(11, DType.U32.tag());
+ assertEquals(12, DType.U64.tag());
+ assertEquals(13, DType.F16.tag());
+ assertEquals(14, DType.values().length);
}
@Test
@@ -173,7 +216,7 @@ public void aDTypeTagPastTheTableIsRejectedByName() {
RuntimeException exception =
assertThrows(RuntimeException.class, () -> NativeExpression.literalNull((byte) DType.values().length));
assertTrue(
- exception.getMessage().contains("unknown null dtype tag: 9"),
+ exception.getMessage().contains("unknown null dtype tag: 14"),
() -> "unexpected message: " + exception.getMessage());
}
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..14567eef54f 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
@@ -10,6 +10,8 @@
import dev.vortex.jni.NativeLoader;
import java.math.BigInteger;
+import java.nio.ByteBuffer;
+import java.nio.ByteOrder;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.Test;
@@ -55,4 +57,148 @@ public void mergeComposes() {
// Merging zero expressions is valid and yields an empty struct.
assertNotNull(Expression.merge());
}
+
+ @Test
+ public void literalTimeAcceptsEveryUnitExceptDays() {
+ for (Expression.TimeUnit unit : new Expression.TimeUnit[] {
+ Expression.TimeUnit.NANOSECONDS,
+ Expression.TimeUnit.MICROSECONDS,
+ Expression.TimeUnit.MILLISECONDS,
+ Expression.TimeUnit.SECONDS
+ }) {
+ assertNotNull(Expression.literalTime(3_600L, unit), unit::name);
+ assertNotNull(Expression.nullLiteralTime(unit), unit::name);
+ }
+ RuntimeException exception =
+ assertThrows(RuntimeException.class, () -> Expression.literalTime(0L, Expression.TimeUnit.DAYS));
+ assertTrue(
+ exception.getMessage().contains("Time type does not support time unit"),
+ () -> "unexpected message: " + exception.getMessage());
+ }
+
+ @Test
+ public void literalTimeRejectsSecondsOutsideI32() {
+ RuntimeException exception = assertThrows(
+ RuntimeException.class,
+ () -> Expression.literalTime((long) Integer.MAX_VALUE + 1, Expression.TimeUnit.SECONDS));
+ assertTrue(
+ exception.getMessage().contains("does not fit in i32"),
+ () -> "unexpected message: " + exception.getMessage());
+ }
+
+ @Test
+ public void literalGeometryBuildsFromWkbPoint() {
+ Expression point = Expression.literalGeometry(wkbPoint(1.0, 2.0));
+ assertNotNull(Expression.binary(Expression.BinaryOp.EQ, Expression.column("geom"), point));
+ }
+
+ @Test
+ public void spatialFunctionsComposeWithGeometryLiterals() {
+ Expression point = Expression.literalGeometry(wkbPoint(1.0, 2.0));
+ Expression intersects =
+ Expression.spatial(Expression.SpatialFunction.INTERSECTS, Expression.column("geom"), point);
+ Expression distance = Expression.spatial(Expression.SpatialFunction.DISTANCE, Expression.column("geom"), point);
+ assertNotNull(Expression.and(
+ intersects, Expression.binary(Expression.BinaryOp.LT, distance, Expression.literal(5.0))));
+ }
+
+ @Test
+ public void spatialRejectsTheWrongArity() {
+ assertThrows(
+ IllegalArgumentException.class,
+ () -> Expression.spatial(
+ Expression.SpatialFunction.AREA, Expression.column("a"), Expression.column("b")));
+ }
+
+ @Test
+ public void literalGeometryRejectsMalformedWkb() {
+ assertThrows(RuntimeException.class, () -> Expression.literalGeometry(new byte[] {1, 2, 3}));
+ }
+
+ @Test
+ public void unsignedLiteralsCoverTheirFullRange() {
+ assertNotNull(Expression.literalU8(0));
+ assertNotNull(Expression.literalU8(255));
+ assertNotNull(Expression.literalU16(65_535));
+ assertNotNull(Expression.literalU32(0xFFFF_FFFFL));
+ assertNotNull(Expression.literalU64(Long.parseUnsignedLong("18446744073709551615")));
+ assertNotNull(Expression.literalU64(BigInteger.ONE.shiftLeft(64).subtract(BigInteger.ONE)));
+ }
+
+ @Test
+ public void unsignedLiteralsRejectOutOfRangeValues() {
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU8(256));
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU8(-1));
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU16(65_536));
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU32(1L << 32));
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU64(BigInteger.ONE.shiftLeft(64)));
+ assertThrows(IllegalArgumentException.class, () -> Expression.literalU64(BigInteger.valueOf(-1)));
+ }
+
+ @Test
+ public void literalF16Builds() {
+ assertNotNull(Expression.binary(Expression.BinaryOp.LT, Expression.column("h"), Expression.literalF16(1.5f)));
+ }
+
+ @Test
+ public void listLiteralsCastElementsToTheElementType() {
+ Expression i32 = Expression.nullLiteral(Expression.DType.I32);
+ // Long elements are cast down to the i32 element type.
+ assertNotNull(Expression.literalList(i32, false, Expression.literal(1L), Expression.literal(2L)));
+ assertNotNull(
+ Expression.literalList(i32, true, Expression.literal(1), Expression.nullLiteral(Expression.DType.I32)));
+ assertNotNull(Expression.literalList(i32, false));
+ assertNotNull(Expression.nullLiteralList(i32, false));
+ }
+
+ @Test
+ public void listLiteralsRejectNonLiteralElements() {
+ Expression i32 = Expression.nullLiteral(Expression.DType.I32);
+ RuntimeException exception =
+ assertThrows(RuntimeException.class, () -> Expression.literalList(i32, false, Expression.column("a")));
+ assertTrue(
+ exception.getMessage().contains("must be a literal expression"),
+ () -> "unexpected message: " + exception.getMessage());
+ }
+
+ @Test
+ public void fixedSizeListLiteralsBuild() {
+ Expression u8 = Expression.nullLiteral(Expression.DType.U8);
+ assertNotNull(Expression.literalFixedSizeList(u8, false, Expression.literalU8(1), Expression.literalU8(2)));
+ assertNotNull(Expression.nullLiteralFixedSizeList(u8, false, 2));
+ }
+
+ @Test
+ public void structLiteralsBuild() {
+ String[] names = {"a", "b"};
+ assertNotNull(
+ Expression.literalStruct(names, new Expression[] {Expression.literal(1), Expression.literal("x")}));
+ assertNotNull(Expression.nullLiteralStruct(names, new Expression[] {
+ Expression.nullLiteral(Expression.DType.I32), Expression.nullLiteral(Expression.DType.UTF8)
+ }));
+ }
+
+ @Test
+ public void mapLiteralsBuild() {
+ Expression keyType = Expression.nullLiteral(Expression.DType.UTF8);
+ Expression valueType = Expression.nullLiteral(Expression.DType.I64);
+ assertNotNull(Expression.literalMap(
+ keyType,
+ valueType,
+ true,
+ new Expression[] {Expression.literal("a"), Expression.literal("b")},
+ new Expression[] {Expression.literal(1L), Expression.nullLiteral(Expression.DType.I64)}));
+ assertNotNull(Expression.nullLiteralMap(keyType, valueType, false));
+ }
+
+ /** Little-endian WKB for {@code POINT(x y)}. */
+ private static byte[] wkbPoint(double x, double y) {
+ return ByteBuffer.allocate(21)
+ .order(ByteOrder.LITTLE_ENDIAN)
+ .put((byte) 1)
+ .putInt(1)
+ .putDouble(x)
+ .putDouble(y)
+ .array();
+ }
}
diff --git a/vortex-jni/src/expression.rs b/vortex-jni/src/expression.rs
index 5f272a24149..66c5c73e382 100644
--- a/vortex-jni/src/expression.rs
+++ b/vortex-jni/src/expression.rs
@@ -28,9 +28,13 @@ use vortex::dtype::BigCast;
use vortex::dtype::DType;
use vortex::dtype::DecimalDType;
use vortex::dtype::FieldName;
+use vortex::dtype::FieldNames;
+use vortex::dtype::MapDType;
use vortex::dtype::Nullability;
use vortex::dtype::PType;
+use vortex::dtype::StructFields;
use vortex::dtype::extension::ExtDType;
+use vortex::dtype::half::f16;
use vortex::encodings::uuid::Uuid;
use vortex::encodings::uuid::UuidMetadata;
use vortex::error::vortex_err;
@@ -48,20 +52,34 @@ use vortex::expr::pack;
use vortex::expr::root;
use vortex::expr::select;
use vortex::extension::datetime::Date;
+use vortex::extension::datetime::Time;
use vortex::extension::datetime::TimeUnit;
use vortex::extension::datetime::Timestamp;
use vortex::layout::layouts::row_idx::row_idx;
use vortex::scalar::DecimalValue;
use vortex::scalar::Scalar;
use vortex::scalar::ScalarValue;
+use vortex::scalar_fn::EmptyOptions;
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::like::Like;
use vortex::scalar_fn::fns::like::LikeOptions;
+use vortex::scalar_fn::fns::literal::Literal;
use vortex::scalar_fn::fns::merge::DuplicateHandling;
use vortex::scalar_fn::fns::operators::Operator;
+use vortex_arrow::ArrowSession;
+use vortex_spatial::extension::native_geometry_scalar_from_wkb;
+use vortex_spatial::scalar_fn::area::SpatialArea;
+use vortex_spatial::scalar_fn::collect::SpatialCollect;
+use vortex_spatial::scalar_fn::contains::SpatialContains;
+use vortex_spatial::scalar_fn::convex_hull::SpatialConvexHull;
+use vortex_spatial::scalar_fn::distance::SpatialDistance;
+use vortex_spatial::scalar_fn::envelope::SpatialEnvelope;
+use vortex_spatial::scalar_fn::intersects::SpatialIntersects;
+use vortex_spatial::scalar_fn::length::SpatialLength;
+use vortex_spatial::scalar_fn::make_line::SpatialMakeLine;
use crate::errors::JNIError;
use crate::errors::try_or_throw;
@@ -278,6 +296,36 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_binary(
})
}
+/// Build a spatial function expression over `operands`.
+///
+/// `function` selects the function; see `dev.vortex.api.Expression.SpatialFunction` on the Java
+/// side for the source of truth.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_spatial(
+ mut env: EnvUnowned,
+ _class: JClass,
+ function: jbyte,
+ operands: JLongArray,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ let operands = collect_operands(env, &operands)?;
+ // `try_new_expr` rejects an operand count that does not match the function's arity.
+ let expr = match function {
+ 0 => SpatialArea.try_new_expr(EmptyOptions, operands),
+ 1 => SpatialCollect.try_new_expr(EmptyOptions, operands),
+ 2 => SpatialContains.try_new_expr(EmptyOptions, operands),
+ 3 => SpatialConvexHull.try_new_expr(EmptyOptions, operands),
+ 4 => SpatialDistance.try_new_expr(EmptyOptions, operands),
+ 5 => SpatialEnvelope.try_new_expr(EmptyOptions, operands),
+ 6 => SpatialIntersects.try_new_expr(EmptyOptions, operands),
+ 7 => SpatialLength.try_new_expr(EmptyOptions, operands),
+ 8 => SpatialMakeLine.try_new_expr(EmptyOptions, operands),
+ other => throw_runtime!("unknown spatial function code: {other}"),
+ }?;
+ Ok(into_raw(expr))
+ })
+}
+
#[unsafe(no_mangle)]
pub extern "system" fn Java_dev_vortex_jni_NativeExpression_not(
_env: EnvUnowned,
@@ -398,6 +446,12 @@ literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalI8, jbyte, i8);
literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalI16, jshort, i16);
literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalI32, jint, i32);
literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalI64, jlong, i64);
+// Java has no unsigned integers, so unsigned literals arrive as the signed type of the same
+// width and are reinterpreted bit-for-bit.
+literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalU8, jbyte, u8);
+literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalU16, jshort, u16);
+literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalU32, jint, u32);
+literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalU64, jlong, u64);
literal_primitive!(Java_dev_vortex_jni_NativeExpression_literalF32, jfloat, f32);
literal_primitive!(
Java_dev_vortex_jni_NativeExpression_literalF64,
@@ -405,6 +459,23 @@ literal_primitive!(
f64
);
+/// Build a half-precision float literal, rounding `value` to the nearest `f16`.
+///
+/// Java has no half-precision type (before `Float.floatToFloat16` in Java 20), so the value
+/// arrives as a `float`.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalF16(
+ _env: EnvUnowned,
+ _class: JClass,
+ value: jfloat,
+ is_null_flag: jboolean,
+) -> jlong {
+ if is_null_flag {
+ return into_raw(lit(Scalar::null_native::()));
+ }
+ into_raw(lit(f16::from_f32(value)))
+}
+
#[unsafe(no_mangle)]
pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalString(
mut env: EnvUnowned,
@@ -596,6 +667,64 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalTimestamp(
})
}
+/// Build a time-of-day literal. `value` is the number of `unit` units since midnight.
+///
+/// Seconds and milliseconds are stored as `i32`, microseconds and nanoseconds as `i64`; days are
+/// rejected by [`Time::try_new`].
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalTime(
+ mut env: EnvUnowned,
+ _class: JClass,
+ value: jlong,
+ time_unit_tag: jbyte,
+ is_null_flag: jboolean,
+) -> jlong {
+ try_or_throw(&mut env, |_| {
+ let unit = parse_time_unit(time_unit_tag)?;
+ let nullability = if is_null_flag {
+ Nullability::Nullable
+ } else {
+ Nullability::NonNullable
+ };
+ let ext = Time::try_new(unit, nullability)?;
+ let dtype = DType::Extension(ext.erased());
+ if is_null_flag {
+ return Ok(into_raw(lit(Scalar::null(dtype))));
+ }
+ let storage_value = match unit {
+ TimeUnit::Seconds | TimeUnit::Milliseconds => ScalarValue::from(
+ i32::try_from(value)
+ .map_err(|_| vortex_err!("time value does not fit in i32 {unit}: {value}"))?,
+ ),
+ _ => ScalarValue::from(value),
+ };
+ Ok(into_raw(lit(Scalar::try_new(dtype, Some(storage_value))?)))
+ })
+}
+
+/// Build a geometry literal from its OGC Well-Known Binary (WKB) encoding.
+///
+/// The value is decoded into the native geometry extension type matching its geometry kind
+/// (`Point`, `LineString`, `Polygon`, `MultiPoint`, `MultiLineString` or `MultiPolygon`, all XY
+/// with no CRS), which is the form the spatial scalar functions and pruning rules operate on.
+/// Geometry collections and malformed WKB are rejected.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalGeometry(
+ mut env: EnvUnowned,
+ _class: JClass,
+ wkb: JByteArray,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ if wkb.is_null() {
+ throw_runtime!("geometry literal WKB bytes must not be null");
+ }
+ let bytes = env.convert_byte_array(&wkb)?;
+ let scalar = native_geometry_scalar_from_wkb(&bytes, &ArrowSession::default())?
+ .ok_or_else(|| vortex_err!("unsupported WKB geometry type for a geometry literal"))?;
+ Ok(into_raw(lit(scalar)))
+ })
+}
+
/// Number of bytes in a UUID's big-endian representation.
const UUID_BYTE_LEN: usize = 16;
@@ -666,6 +795,213 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalUuid(
})
}
+/// The scalar held by a literal expression, or an error naming `role` if `ptr` is not a literal.
+///
+/// SAFETY: `ptr` must satisfy the contract of [`expr_ref`].
+unsafe fn literal_scalar<'a>(ptr: jlong, role: &str) -> Result<&'a Scalar, JNIError> {
+ let expr = unsafe { expr_ref(ptr) };
+ expr.as_opt::()
+ .ok_or_else(|| vortex_err!("{role} must be a literal expression, got {expr}").into())
+}
+
+/// The dtype of the literal `ptr`, used as a type prototype, with `nullability` applied.
+///
+/// SAFETY: `ptr` must satisfy the contract of [`expr_ref`].
+unsafe fn prototype_dtype(
+ ptr: jlong,
+ role: &str,
+ nullability: Nullability,
+) -> Result {
+ Ok(unsafe { literal_scalar(ptr, role) }?
+ .dtype()
+ .with_nullability(nullability))
+}
+
+/// Read the literal expressions in `pointers` and cast each one to `dtype`.
+fn cast_literals(
+ env: &mut jni::Env,
+ pointers: &JLongArray,
+ role: &str,
+ dtype: &DType,
+) -> Result, JNIError> {
+ let ptrs = unsafe { pointers.get_elements(env, ReleaseMode::NoCopyBack) }?;
+ ptrs.iter()
+ .map(|ptr| -> Result {
+ Ok(unsafe { literal_scalar(*ptr, role) }?.cast(dtype)?)
+ })
+ .collect()
+}
+
+/// Build a variable-length list literal.
+///
+/// Every element must be a literal expression and is cast to the element dtype: the dtype of the
+/// `element_type` literal (typically a typed null) with nullability `elements_nullable`. With
+/// `is_null_flag` the `elements` are ignored and a null list of that element dtype is produced.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalList(
+ mut env: EnvUnowned,
+ _class: JClass,
+ elements: JLongArray,
+ element_type: jlong,
+ elements_nullable: jboolean,
+ is_null_flag: jboolean,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ let element_dtype = unsafe {
+ prototype_dtype(element_type, "list element type", elements_nullable.into())
+ }?;
+ if is_null_flag {
+ return Ok(into_raw(lit(Scalar::null(DType::List(
+ Arc::new(element_dtype),
+ Nullability::Nullable,
+ )))));
+ }
+ let children = cast_literals(env, &elements, "list element", &element_dtype)?;
+ Ok(into_raw(lit(Scalar::list(
+ element_dtype,
+ children,
+ Nullability::NonNullable,
+ ))))
+ })
+}
+
+/// Build a fixed-size list literal of `size` elements.
+///
+/// Elements and `element_type` follow [`Java_dev_vortex_jni_NativeExpression_literalList`]. A
+/// non-null literal must have exactly `size` elements.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalFixedSizeList(
+ mut env: EnvUnowned,
+ _class: JClass,
+ elements: JLongArray,
+ element_type: jlong,
+ elements_nullable: jboolean,
+ size: jint,
+ is_null_flag: jboolean,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ let size = u32::try_from(size)
+ .map_err(|_| vortex_err!("fixed-size list size must not be negative: {size}"))?;
+ let element_dtype = unsafe {
+ prototype_dtype(
+ element_type,
+ "fixed-size list element type",
+ elements_nullable.into(),
+ )
+ }?;
+ if is_null_flag {
+ return Ok(into_raw(lit(Scalar::null(DType::FixedSizeList(
+ Arc::new(element_dtype),
+ size,
+ Nullability::Nullable,
+ )))));
+ }
+ let children = cast_literals(env, &elements, "fixed-size list element", &element_dtype)?;
+ if children.len() != size as usize {
+ throw_runtime!(
+ "fixed-size list literal of size {size} has {} elements",
+ children.len()
+ );
+ }
+ Ok(into_raw(lit(Scalar::fixed_size_list(
+ element_dtype,
+ children,
+ Nullability::NonNullable,
+ ))))
+ })
+}
+
+/// Build a struct literal from named literal fields.
+///
+/// The struct's field dtypes are the dtypes of the `fields` literals. With `is_null_flag` the
+/// `fields` only supply those dtypes (typically as typed nulls) and a null struct is produced.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalStruct(
+ mut env: EnvUnowned,
+ _class: JClass,
+ field_names: JObjectArray,
+ fields: JLongArray,
+ is_null_flag: jboolean,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ let count = field_names.len(env)?;
+ let ptrs = unsafe { fields.get_elements(env, ReleaseMode::NoCopyBack) }?;
+ if ptrs.len() != count {
+ throw_runtime!(
+ "struct literal has {count} field names but {} fields",
+ ptrs.len()
+ );
+ }
+ let mut names: Vec = Vec::with_capacity(count);
+ let mut children: Vec = Vec::with_capacity(count);
+ for (idx, ptr) in ptrs.iter().enumerate() {
+ let obj = field_names.get_element(env, idx)?;
+ let name = env.cast_local::(obj)?;
+ names.push(name.try_to_string(env)?.into());
+ children.push(unsafe { literal_scalar(*ptr, "struct field") }?.clone());
+ }
+
+ let struct_fields = StructFields::new(
+ FieldNames::from(names),
+ children.iter().map(|child| child.dtype().clone()).collect(),
+ );
+ if is_null_flag {
+ return Ok(into_raw(lit(Scalar::null(DType::Struct(
+ struct_fields,
+ Nullability::Nullable,
+ )))));
+ }
+ Ok(into_raw(lit(Scalar::struct_(
+ DType::Struct(struct_fields, Nullability::NonNullable),
+ children,
+ ))))
+ })
+}
+
+/// Build a map literal from parallel arrays of literal keys and values.
+///
+/// Keys are cast to the non-nullable dtype of `key_type`, and values to the dtype of `value_type`
+/// with nullability `values_nullable`. With `is_null_flag` the entries are ignored and a null map
+/// of that type is produced.
+#[unsafe(no_mangle)]
+pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalMap(
+ mut env: EnvUnowned,
+ _class: JClass,
+ keys: JLongArray,
+ values: JLongArray,
+ key_type: jlong,
+ value_type: jlong,
+ values_nullable: jboolean,
+ is_null_flag: jboolean,
+) -> jlong {
+ try_or_throw(&mut env, |env| {
+ let key_dtype =
+ unsafe { prototype_dtype(key_type, "map key type", Nullability::NonNullable) }?;
+ let value_dtype =
+ unsafe { prototype_dtype(value_type, "map value type", values_nullable.into()) }?;
+ let map_dtype = MapDType::try_new(key_dtype.clone(), value_dtype.clone(), false)?;
+ if is_null_flag {
+ return Ok(into_raw(lit(Scalar::null(DType::Map(
+ map_dtype,
+ Nullability::Nullable,
+ )))));
+ }
+ let keys = cast_literals(env, &keys, "map key", &key_dtype)?;
+ let values = cast_literals(env, &values, "map value", &value_dtype)?;
+ if keys.len() != values.len() {
+ throw_runtime!(
+ "map literal has {} keys but {} values",
+ keys.len(),
+ values.len()
+ );
+ }
+ Ok(into_raw(lit(Scalar::try_map(
+ DType::Map(map_dtype, Nullability::NonNullable),
+ keys.into_iter().zip(values),
+ )?)))
+ })
+}
+
/// Build a typed null literal whose nullable dtype is selected by `dtype_tag`.
///
/// Tag values intentionally do not overlap with [`parse_time_unit`].
@@ -687,6 +1023,11 @@ pub extern "system" fn Java_dev_vortex_jni_NativeExpression_literalNull(
6 => DType::Primitive(PType::F64, Nullability::Nullable),
7 => DType::Utf8(Nullability::Nullable),
8 => DType::Binary(Nullability::Nullable),
+ 9 => DType::Primitive(PType::U8, Nullability::Nullable),
+ 10 => DType::Primitive(PType::U16, Nullability::Nullable),
+ 11 => DType::Primitive(PType::U32, Nullability::Nullable),
+ 12 => DType::Primitive(PType::U64, Nullability::Nullable),
+ 13 => DType::Primitive(PType::F16, Nullability::Nullable),
other => throw_runtime!("unknown null dtype tag: {other}"),
};
Ok(into_raw(lit(Scalar::null(dtype))))
diff --git a/vortex-python/python/vortex/_lib/expr.pyi b/vortex-python/python/vortex/_lib/expr.pyi
index 6cd2b58bfd5..8463f113574 100644
--- a/vortex-python/python/vortex/_lib/expr.pyi
+++ b/vortex-python/python/vortex/_lib/expr.pyi
@@ -2,15 +2,17 @@
# SPDX-FileCopyrightText: Copyright the Vortex contributors
from collections.abc import Iterable, Mapping, Sequence
-from datetime import date, datetime
+from datetime import date, datetime, time
+from decimal import Decimal
from typing import Literal, TypeAlias, final
+from uuid import UUID
from typing_extensions import override
from .dtype import DType
from .scalar import ScalarPyType
-IntoExpr: TypeAlias = Expr | bool | int | float | str | bytes | date | datetime | None
+IntoExpr: TypeAlias = Expr | bool | int | float | str | bytes | Decimal | date | datetime | time | UUID | None
"""A value accepted anywhere an expression is expected. Non-``Expr`` values become literals."""
VariantPath: TypeAlias = str | int | Sequence[str | int]
diff --git a/vortex-python/src/expr/mod.rs b/vortex-python/src/expr/mod.rs
index 8d76d20b051..7355eebcd0b 100644
--- a/vortex-python/src/expr/mod.rs
+++ b/vortex-python/src/expr/mod.rs
@@ -141,8 +141,9 @@ impl PyExpr {
/// A Python value that can be coerced into an [`Expression`].
///
/// Accepts an existing [`PyExpr`], or any Python value convertible to a Vortex scalar (including
-/// `None`, `bool`, `int`, `float`, `str`, `bytes`, `list`, `dict`, and `vortex.Scalar`), which is
-/// wrapped in a literal expression.
+/// `None`, `bool`, `int`, `float`, `str`, `bytes`, `list`, `dict`, `decimal.Decimal`,
+/// `datetime.date`, `datetime.datetime`, `datetime.time`, `uuid.UUID`, and `vortex.Scalar`),
+/// which is wrapped in a literal expression.
pub struct PyIntoExpr(Expression);
impl PyIntoExpr {
diff --git a/vortex-python/src/scalar/factory.rs b/vortex-python/src/scalar/factory.rs
index 13e46b048c9..54ec36228ac 100644
--- a/vortex-python/src/scalar/factory.rs
+++ b/vortex-python/src/scalar/factory.rs
@@ -5,6 +5,7 @@ use std::sync::Arc;
use itertools::Itertools;
use pyo3::exceptions::PyValueError;
+use pyo3::intern;
use pyo3::prelude::*;
use pyo3::types::PyBool;
use pyo3::types::PyBytes;
@@ -13,19 +14,76 @@ use pyo3::types::PyFloat;
use pyo3::types::PyInt;
use pyo3::types::PyList;
use pyo3::types::PyString;
+use vortex::dtype::BigCast;
use vortex::dtype::DType;
+use vortex::dtype::DecimalDType;
use vortex::dtype::FieldName;
use vortex::dtype::FieldNames;
+use vortex::dtype::MAX_PRECISION;
use vortex::dtype::Nullability;
+use vortex::dtype::PType;
use vortex::dtype::StructFields;
+use vortex::dtype::extension::ExtDType;
+use vortex::dtype::i256;
+use vortex::encodings::uuid::Uuid;
+use vortex::encodings::uuid::UuidMetadata;
+use vortex::extension::datetime::Date;
+use vortex::extension::datetime::Time;
+use vortex::extension::datetime::TimeUnit;
+use vortex::extension::datetime::Timestamp;
use vortex::scalar::DecimalValue;
use vortex::scalar::Scalar;
+use vortex::scalar::ScalarValue;
use crate::dtype::PyDType;
use crate::error::PyVortexResult;
use crate::scalar::PyScalar;
use crate::scalar::bool;
+/// Construct a Vortex scalar from a Python value.
+///
+/// Parameters
+/// ----------
+/// value : :class:`object`
+/// The value to convert. Supported types are :obj:`None`, :class:`bool`, :class:`int`,
+/// :class:`float`, :class:`str`, :class:`bytes`, :class:`list`, :class:`dict` (as a struct),
+/// :class:`decimal.Decimal`, :class:`datetime.date`, :class:`datetime.datetime`,
+/// :class:`datetime.time`, :class:`uuid.UUID`, and :class:`vortex.Scalar`.
+/// dtype : :class:`vortex.DType`, optional
+/// The data type of the result. When given, it also guides the conversion: the unit of a
+/// date, time or timestamp type, the scale of a decimal type, and the element and field types
+/// of list and struct types. Without it, :class:`int` becomes a 64-bit integer,
+/// :class:`float` a 64-bit float, :class:`decimal.Decimal` a decimal with the value's own
+/// precision and scale, dates use days, and times and timestamps use microseconds.
+///
+/// Returns
+/// -------
+/// :class:`vortex.Scalar`
+///
+/// Raises
+/// ------
+/// ValueError
+/// If the value cannot be represented exactly, for example a decimal or timestamp that
+/// would lose precision in the requested scale or unit, a non-finite decimal, or a
+/// timezone that has no IANA name.
+///
+/// Notes
+/// -----
+/// Comparisons in expressions require both sides to have the same type, so a literal compared
+/// against a column should match the column's type, for example
+/// ``vx.scalar(value, dtype=vx.timestamp("ns"))`` for a nanosecond timestamp column.
+///
+/// Examples
+/// --------
+///
+/// ```python
+/// >>> import datetime
+/// >>> import vortex as vx
+/// >>> vx.scalar(datetime.date(1970, 1, 2)).as_py()
+/// 1
+/// >>> vx.scalar(datetime.time(0, 0, 1), dtype=vx.time("ms")).as_py()
+/// 1000
+/// ```
#[pyfunction(name = "scalar")]
#[pyo3(signature = (value, *, dtype=None))]
pub fn scalar<'py>(
@@ -77,6 +135,9 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes
// decimal
if let Some(decimal_dtype) = dtype.and_then(|d| d.as_decimal_opt()) {
+ if is_decimal(value)? {
+ return Ok(decimal_scalar(value, Some(decimal_dtype))?);
+ }
let value = if let Ok(v) = value.extract::() {
DecimalValue::I8(v)
} else if let Ok(v) = value.extract::() {
@@ -150,10 +211,15 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes
)));
}
+ // Each field is cast to its dtype, since `Scalar::struct_` requires every child to
+ // match it exactly.
let children: Vec = dict
.values()
.into_iter()
- .map(|item| scalar_helper_inner(&item, None))
+ .zip(dtype.fields())
+ .map(|(item, field_dtype)| {
+ scalar_helper(&item, Some(&field_dtype)).map_err(PyErr::from)
+ })
.try_collect()?;
return Ok(Scalar::struct_(
DType::Struct(dtype.clone(), *nullability),
@@ -178,15 +244,17 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes
if let Ok(list) = value.cast::() {
if let Some(DType::List(element_dtype, ..)) = dtype {
+ // Each element is cast to the element dtype, since `Scalar::list` requires every
+ // child to match it exactly.
let elements = list
.iter()
- .map(|e| scalar_helper_inner(&e, Some(element_dtype)))
+ .map(|e| scalar_helper(&e, Some(element_dtype)).map_err(PyErr::from))
.try_collect()?;
- Scalar::list(
+ 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.
@@ -212,8 +280,405 @@ fn scalar_helper_inner(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyRes
}
}
+ // Standard library types, checked after the built-in types to keep those conversions cheap.
+ let py = value.py();
+
+ // decimal.Decimal without a decimal dtype hint
+ if is_decimal(value)? {
+ return Ok(decimal_scalar(value, None)?);
+ }
+
+ // datetime.datetime, checked before datetime.date because it is a subclass of it.
+ let datetime_module = py.import(intern!(py, "datetime"))?;
+ if value.is_instance(&datetime_module.getattr(intern!(py, "datetime"))?)? {
+ return Ok(timestamp_scalar(value, dtype)?);
+ }
+
+ // datetime.date
+ if value.is_instance(&datetime_module.getattr(intern!(py, "date"))?)? {
+ return Ok(date_scalar(value, dtype)?);
+ }
+
+ // datetime.time
+ if value.is_instance(&datetime_module.getattr(intern!(py, "time"))?)? {
+ return Ok(time_scalar(value, dtype)?);
+ }
+
+ // uuid.UUID
+ let uuid_type = py
+ .import(intern!(py, "uuid"))?
+ .getattr(intern!(py, "UUID"))?;
+ if value.is_instance(&uuid_type)? {
+ let bytes: Vec = value.getattr(intern!(py, "bytes"))?.extract()?;
+ return Ok(uuid_scalar(&bytes)?);
+ }
+
Err(pyo3::exceptions::PyTypeError::new_err(format!(
"Cannot convert Python object to Vortex scalar: {}",
value.get_type()
)))
}
+
+/// Proleptic Gregorian ordinal of 1970-01-01, as returned by `date.toordinal()`.
+const UNIX_EPOCH_ORDINAL: i64 = 719_163;
+const MICROS_PER_SECOND: i64 = 1_000_000;
+const MICROS_PER_DAY: i64 = 86_400 * MICROS_PER_SECOND;
+/// Number of bytes in a UUID.
+const UUID_BYTE_LEN: usize = 16;
+
+fn is_decimal(value: &Bound<'_, PyAny>) -> PyResult {
+ let py = value.py();
+ let decimal_type = py
+ .import(intern!(py, "decimal"))?
+ .getattr(intern!(py, "Decimal"))?;
+ value.is_instance(&decimal_type)
+}
+
+/// Convert a count of microseconds into `unit`, failing rather than silently truncating.
+fn micros_to_unit(micros: i64, unit: TimeUnit) -> PyResult {
+ let per_unit = match unit {
+ TimeUnit::Nanoseconds => {
+ return micros.checked_mul(1_000).ok_or_else(|| {
+ PyValueError::new_err(format!("{micros}us overflows i64 nanoseconds"))
+ });
+ }
+ TimeUnit::Microseconds => return Ok(micros),
+ TimeUnit::Milliseconds => 1_000,
+ TimeUnit::Seconds => MICROS_PER_SECOND,
+ TimeUnit::Days => MICROS_PER_DAY,
+ };
+ if micros % per_unit != 0 {
+ return Err(PyValueError::new_err(format!(
+ "{micros}us cannot be represented in {unit} without losing precision"
+ )));
+ }
+ Ok(micros / per_unit)
+}
+
+/// Microseconds since midnight of a `datetime.time` or `datetime.datetime`.
+fn time_of_day_micros(value: &Bound<'_, PyAny>) -> PyResult {
+ let py = value.py();
+ let hour: i64 = value.getattr(intern!(py, "hour"))?.extract()?;
+ let minute: i64 = value.getattr(intern!(py, "minute"))?.extract()?;
+ let second: i64 = value.getattr(intern!(py, "second"))?.extract()?;
+ let microsecond: i64 = value.getattr(intern!(py, "microsecond"))?.extract()?;
+ Ok(((hour * 60 + minute) * 60 + second) * MICROS_PER_SECOND + microsecond)
+}
+
+/// Days since the Unix epoch of a `datetime.date` or `datetime.datetime`.
+fn epoch_days(value: &Bound<'_, PyAny>) -> PyResult {
+ let ordinal: i64 = value
+ .call_method0(intern!(value.py(), "toordinal"))?
+ .extract()?;
+ Ok(ordinal - UNIX_EPOCH_ORDINAL)
+}
+
+/// The Vortex timezone name of a `tzinfo`.
+///
+/// Vortex timestamps carry an IANA zone name, so only `zoneinfo.ZoneInfo` and pytz zones and
+/// `datetime.timezone.utc` are accepted; other fixed-offset zones have no IANA name.
+fn timezone_name(tzinfo: &Bound<'_, PyAny>) -> PyResult {
+ let py = tzinfo.py();
+ // `zoneinfo.ZoneInfo` names its zone `key`; pytz zones name it `zone`.
+ for attr in [intern!(py, "key"), intern!(py, "zone")] {
+ if let Ok(name) = tzinfo.getattr(attr)
+ && let Ok(Some(name)) = name.extract::