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::>() + { + return Ok(name); + } + } + let utc = py + .import(intern!(py, "datetime"))? + .getattr(intern!(py, "timezone"))? + .getattr(intern!(py, "utc"))?; + if tzinfo.eq(utc)? { + return Ok("UTC".to_string()); + } + Err(PyValueError::new_err(format!( + "Unsupported timezone {}: use a zoneinfo.ZoneInfo, a pytz zone or datetime.timezone.utc", + tzinfo.repr()? + ))) +} + +/// Convert a `datetime.date` into a Vortex `Date` scalar. +/// +/// The unit is taken from `dtype` when it is a `Date` dtype, and otherwise defaults to days. +fn date_scalar(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyVortexResult { + let unit = dtype + .and_then(DType::as_extension_opt) + .and_then(|ext| ext.metadata_opt::()) + .copied() + .unwrap_or(TimeUnit::Days); + let ext = Date::try_new(unit, Nullability::NonNullable)?; + let days = epoch_days(value)?; + let storage = match unit { + TimeUnit::Days => ScalarValue::from( + i32::try_from(days) + .map_err(|_| PyValueError::new_err(format!("{days} days does not fit in i32")))?, + ), + _ => ScalarValue::from(micros_to_unit(days * MICROS_PER_DAY, unit)?), + }; + Ok(Scalar::try_new( + DType::Extension(ext.erased()), + Some(storage), + )?) +} + +/// Convert a `datetime.datetime` into a Vortex `Timestamp` scalar. +/// +/// The unit and timezone are taken from `dtype` when it is a `Timestamp` dtype, and otherwise +/// default to microseconds and the value's own timezone. Naive values are stored as their +/// wall-clock time; timezone-aware values are stored as the UTC instant. A naive value cannot +/// be converted to a timezone-aware dtype, nor an aware value to a naive dtype. +fn timestamp_scalar(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyVortexResult { + let py = value.py(); + let tzinfo = value.getattr(intern!(py, "tzinfo"))?; + let aware = !tzinfo.is_none(); + + let mut micros = epoch_days(value)? * MICROS_PER_DAY + time_of_day_micros(value)?; + if aware { + let offset = value.call_method0(intern!(py, "utcoffset"))?; + let days: i64 = offset.getattr(intern!(py, "days"))?.extract()?; + let seconds: i64 = offset.getattr(intern!(py, "seconds"))?.extract()?; + let microseconds: i64 = offset.getattr(intern!(py, "microseconds"))?.extract()?; + micros -= days * MICROS_PER_DAY + seconds * MICROS_PER_SECOND + microseconds; + } + + let options = dtype + .and_then(DType::as_extension_opt) + .and_then(|ext| ext.metadata_opt::()); + let (unit, tz) = match options { + Some(options) => { + if options.tz.is_some() != aware { + return Err(PyValueError::new_err(format!( + "Cannot convert a {} datetime to a timestamp dtype {} a timezone", + if aware { "timezone-aware" } else { "naive" }, + if aware { "without" } else { "with" }, + )) + .into()); + } + (options.unit, options.tz.clone()) + } + None => { + let tz = if aware { + Some(Arc::from(timezone_name(&tzinfo)?.as_str())) + } else { + None + }; + (TimeUnit::Microseconds, tz) + } + }; + + // `pandas.Timestamp` subclasses `datetime.datetime` and keeps its sub-microsecond part in + // `nanosecond`, which `microsecond` excludes. + let nanosecond: i64 = match value.getattr(intern!(py, "nanosecond")) { + Ok(nanosecond) => nanosecond.extract()?, + Err(_) => 0, + }; + let storage = if unit == TimeUnit::Nanoseconds { + micros_to_unit(micros, unit)? + .checked_add(nanosecond) + .ok_or_else(|| PyValueError::new_err("Timestamp overflows i64 nanoseconds"))? + } else if nanosecond != 0 { + return Err(PyValueError::new_err(format!( + "Timestamp with {nanosecond}ns cannot be represented in {unit} without losing \ + precision; pass dtype=vx.timestamp(\"ns\")" + )) + .into()); + } else { + micros_to_unit(micros, unit)? + }; + + let ext = Timestamp::new_with_tz(unit, tz, Nullability::NonNullable); + Ok(Scalar::try_new( + DType::Extension(ext.erased()), + Some(ScalarValue::from(storage)), + )?) +} + +/// Convert a `decimal.Decimal` into a Vortex decimal scalar. +/// +/// With a decimal dtype hint the value is rescaled to its scale exactly; otherwise the precision +/// and scale are inferred from the value's digits and exponent. Non-finite values, values with +/// more than [`MAX_PRECISION`] digits, and rescaling that would drop non-zero digits are errors. +fn decimal_scalar(value: &Bound<'_, PyAny>, hint: Option<&DecimalDType>) -> PyVortexResult { + let py = value.py(); + let parts = value.call_method0(intern!(py, "as_tuple"))?; + // The exponent is a string ('n', 'N' or 'F') for NaN and infinities. + let Ok(exponent) = parts.getattr(intern!(py, "exponent"))?.extract::() else { + return Err(PyValueError::new_err(format!( + "Cannot convert non-finite decimal {} to a Vortex scalar", + value.str()? + )) + .into()); + }; + let negative = parts.getattr(intern!(py, "sign"))?.extract::()? == 1; + let digits: Vec = parts.getattr(intern!(py, "digits"))?.extract()?; + let digits = match digits.iter().position(|&d| d != 0) { + Some(first) => &digits[first..], + None => &[][..], + }; + + let scale: i64 = hint.map_or_else(|| (-exponent).max(0), |d| d.scale().into()); + // Shift the digits so that they represent the unscaled value at `scale`. + let shift = exponent + scale; + let unscaled_digits: Vec = if digits.is_empty() { + Vec::new() + } else if let Ok(zeros) = usize::try_from(shift) { + if digits.len().saturating_add(zeros) > usize::from(MAX_PRECISION) { + return Err(PyValueError::new_err(format!( + "Decimal {} has more than {MAX_PRECISION} digits at scale {scale}", + value.str()? + )) + .into()); + } + let mut shifted = digits.to_vec(); + shifted.resize(digits.len() + zeros, 0); + shifted + } else { + let keep = digits + .len() + .saturating_sub(usize::try_from(-shift).unwrap_or(usize::MAX)); + if digits[keep..].iter().any(|&d| d != 0) { + return Err(PyValueError::new_err(format!( + "Decimal {} cannot be represented with scale {scale} without losing precision", + value.str()? + )) + .into()); + } + digits[..keep].to_vec() + }; + + let decimal_dtype = match hint { + Some(hint) => { + if unscaled_digits.len() > usize::from(hint.precision()) { + return Err(PyValueError::new_err(format!( + "Decimal {} does not fit in precision {}", + value.str()?, + hint.precision() + )) + .into()); + } + *hint + } + None => { + let precision = i64::try_from(unscaled_digits.len()) + .unwrap_or(i64::MAX) + .max(scale) + .max(1); + DecimalDType::try_new( + u8::try_from(precision).unwrap_or(u8::MAX), + i8::try_from(scale).unwrap_or(i8::MAX), + ) + .map_err(|err| { + PyValueError::new_err(format!( + "Decimal {} cannot be represented as a Vortex decimal: {err}", + value + )) + })? + } + }; + + let ten = i256::from_i128(10); + let mut unscaled = i256::ZERO; + for digit in unscaled_digits { + unscaled = unscaled * ten + i256::from_i128(digit.into()); + } + if negative { + unscaled = -unscaled; + } + + Ok(Scalar::try_new( + DType::Decimal(decimal_dtype, Nullability::NonNullable), + Some(ScalarValue::from(narrowest_decimal_value( + unscaled, + &decimal_dtype, + )?)), + )?) +} + +/// The narrowest [`DecimalValue`] variant able to hold every value of `dtype`'s precision. +fn narrowest_decimal_value(value: i256, dtype: &DecimalDType) -> PyResult { + let overflow = || PyValueError::new_err(format!("Decimal value {value} overflows {dtype}")); + let required_bits = dtype.required_bit_width(); + Ok(if required_bits <= 8 { + DecimalValue::I8(BigCast::from(value).ok_or_else(overflow)?) + } else if required_bits <= 16 { + DecimalValue::I16(BigCast::from(value).ok_or_else(overflow)?) + } else if required_bits <= 32 { + DecimalValue::I32(BigCast::from(value).ok_or_else(overflow)?) + } else if required_bits <= 64 { + DecimalValue::I64(BigCast::from(value).ok_or_else(overflow)?) + } else if required_bits <= 128 { + DecimalValue::I128(value.maybe_i128().ok_or_else(overflow)?) + } else { + DecimalValue::I256(value) + }) +} + +/// Build a UUID extension scalar from its 16 big-endian bytes, matching `uuid.UUID.bytes`. +fn uuid_scalar(bytes: &[u8]) -> PyVortexResult { + let list_size = u32::try_from(UUID_BYTE_LEN).map_err(|_| { + PyValueError::new_err(format!( + "UUID byte length {UUID_BYTE_LEN} does not fit in u32" + )) + })?; + if bytes.len() != UUID_BYTE_LEN { + return Err(PyValueError::new_err(format!( + "UUID must be exactly {UUID_BYTE_LEN} bytes, got {}", + bytes.len() + )) + .into()); + } + let element_dtype = DType::Primitive(PType::U8, Nullability::NonNullable); + let storage = Scalar::fixed_size_list( + element_dtype.clone(), + bytes + .iter() + .map(|&b| Scalar::primitive(b, Nullability::NonNullable)) + .collect(), + Nullability::NonNullable, + ); + let ext = ExtDType::::try_new( + UuidMetadata::default(), + DType::FixedSizeList(Arc::new(element_dtype), list_size, Nullability::NonNullable), + )?; + Ok(Scalar::try_new( + DType::Extension(ext.erased()), + storage.into_value(), + )?) +} + +/// Convert a naive `datetime.time` into a Vortex `Time` scalar. +/// +/// The unit is taken from `dtype` when it is a `Time` dtype, and otherwise defaults to +/// microseconds, the resolution of `datetime.time`. Converting to a coarser unit that would drop +/// a non-zero fraction of a second is an error rather than a silent truncation. +fn time_scalar(value: &Bound<'_, PyAny>, dtype: Option<&DType>) -> PyVortexResult { + let py = value.py(); + if !value.getattr(intern!(py, "tzinfo"))?.is_none() { + return Err(PyValueError::new_err( + "Timezone-aware datetime.time values cannot be converted to a Vortex time scalar", + ) + .into()); + } + let unit = dtype + .and_then(DType::as_extension_opt) + .and_then(|ext| ext.metadata_opt::