diff --git a/java/src/main/java/ai/rapids/cudf/VariantUtils.java b/java/src/main/java/ai/rapids/cudf/VariantUtils.java index 4c8e11d9df3..74186a0f864 100644 --- a/java/src/main/java/ai/rapids/cudf/VariantUtils.java +++ b/java/src/main/java/ai/rapids/cudf/VariantUtils.java @@ -21,7 +21,8 @@ public class VariantUtils { // cpp/include/cudf/io/experimental/variant.hpp and // cpp/src/io/parquet/experimental/variant_extract.cu:is_variant_castable. private static final List SUPPORTED_TYPES = Arrays.asList( - DType.STRING, DType.INT8, DType.INT16, DType.INT32, DType.INT64); + DType.STRING, DType.INT8, DType.INT16, DType.INT32, DType.INT64, + DType.FLOAT32, DType.FLOAT64, DType.BOOL8); private VariantUtils() {} @@ -38,8 +39,11 @@ private static void validateTargetType(DType targetType) { * * @param variantStruct Variant materialization: STRUCT(metadata LIST<UINT8>, * value LIST<UINT8>, optional shredded children...) - * @param path JSONPath-like path accepted by cuDF's Variant extractor. Paths are expected to - * be ASCII object-field paths like {@code x}, {@code $.x}, or {@code $.x.y}. + * @param path JSONPath-like path accepted by cuDF's Variant extractor. Object-field steps use + * dot notation and zero-based array-index steps use bracket notation, for example + * {@code x}, {@code $.x.y}, {@code $[0]}, or {@code $.a[0].b}. Missing fields, + * out-of-bounds indices, and container mismatches produce null rows. Wildcards, + * negative indices, and quoted names inside brackets are not supported. * @return LIST<UINT8> column of raw encoded Variant values */ public static ColumnVector getVariantFieldValue(ColumnView variantStruct, String path) { @@ -50,8 +54,11 @@ public static ColumnVector getVariantFieldValue(ColumnView variantStruct, String /** * Decode raw Variant-encoded value bytes into {@code targetType}. Supported target types are - * {@link DType#STRING}, {@link DType#INT8}, {@link DType#INT16}, {@link DType#INT32}, and - * {@link DType#INT64}. + * {@link DType#STRING}, {@link DType#INT8}, {@link DType#INT16}, {@link DType#INT32}, + * {@link DType#INT64}, {@link DType#FLOAT32}, {@link DType#FLOAT64}, and {@link DType#BOOL8}. + * Decoding requires the encoded physical type to exactly match {@code targetType}; no numeric + * conversions are performed. Input nulls, encoded Variant nulls, and physical-type mismatches + * produce null output rows. */ public static ColumnVector castVariantValue(ColumnView valueBytes, DType targetType) { Objects.requireNonNull(valueBytes, "valueBytes"); @@ -63,7 +70,20 @@ public static ColumnVector castVariantValue(ColumnView valueBytes, DType targetT /** * Extract a Variant field and decode it into {@code targetType} in one native call. * Supported target types are {@link DType#STRING}, {@link DType#INT8}, {@link DType#INT16}, - * {@link DType#INT32}, and {@link DType#INT64}. + * {@link DType#INT32}, {@link DType#INT64}, {@link DType#FLOAT32}, {@link DType#FLOAT64}, and + * {@link DType#BOOL8}. + * Decoding requires the encoded physical type to exactly match {@code targetType}; no numeric + * conversions are performed. Missing fields, input nulls, encoded Variant nulls, and + * physical-type mismatches produce null output rows. + * + * @param variantStruct Variant materialization: STRUCT(metadata LIST<UINT8>, + * value LIST<UINT8>, optional shredded children...) + * @param path JSONPath-like path accepted by cuDF's Variant extractor. Object-field steps use + * dot notation and zero-based array-index steps use bracket notation, for example + * {@code x}, {@code $.x.y}, {@code $[0]}, or {@code $.a[0].b}. Missing fields, + * out-of-bounds indices, and container mismatches produce null rows. Wildcards, + * negative indices, and quoted names inside brackets are not supported. + * @param targetType decoded output type */ public static ColumnVector extractVariantField( ColumnView variantStruct, String path, DType targetType) { diff --git a/java/src/test/java/ai/rapids/cudf/VariantUtilsTest.java b/java/src/test/java/ai/rapids/cudf/VariantUtilsTest.java index d9c9633805e..5402a0a2f12 100644 --- a/java/src/test/java/ai/rapids/cudf/VariantUtilsTest.java +++ b/java/src/test/java/ai/rapids/cudf/VariantUtilsTest.java @@ -155,6 +155,24 @@ private List object(Map> fields, List valueOrde return concat(bytes(header, fields.size()), fieldIds, offsets, encodedValues); } + @SafeVarargs + private static List array(List... values) { + List offsets = new ArrayList<>(values.length + 1); + List encodedValues = new ArrayList<>(); + int offset = 0; + for (List value : values) { + offsets.add(oneByte(offset, "array offset")); + offset += value.size(); + encodedValues.addAll(value); + } + offsets.add(oneByte(offset, "array offset")); + + int header = header(0x03, + ((SMALL_OFFSET_SIZE - 1) & 0x03) + | ((SMALL_CONTAINER & 0x01) << 2)); + return concat(bytes(header, values.length), offsets, encodedValues); + } + private static List sortedFieldNames(Map> fields) { List sortedNames = new ArrayList<>(fields.keySet()); sortedNames.sort(VariantEncoder::compareUtf8Unsigned); @@ -208,6 +226,36 @@ private static List int64(long value) { (int) ((value >>> 56) & 0xff)); } + private static List float32(float value) { + int bits = Float.floatToRawIntBits(value); + return bytes(simple(0x0e), + bits & 0xff, + (bits >>> 8) & 0xff, + (bits >>> 16) & 0xff, + (bits >>> 24) & 0xff); + } + + private static List float64(double value) { + long bits = Double.doubleToRawLongBits(value); + return bytes(simple(0x07), + (int) (bits & 0xff), + (int) ((bits >>> 8) & 0xff), + (int) ((bits >>> 16) & 0xff), + (int) ((bits >>> 24) & 0xff), + (int) ((bits >>> 32) & 0xff), + (int) ((bits >>> 40) & 0xff), + (int) ((bits >>> 48) & 0xff), + (int) ((bits >>> 56) & 0xff)); + } + + private static List bool(boolean value) { + return bytes(simple(value ? 0x01 : 0x02)); + } + + private static List nullValue() { + return bytes(simple(0x00)); + } + private static List string(String value) { byte[] bytes = value.getBytes(StandardCharsets.UTF_8); if (bytes.length >= (1 << 6)) { @@ -304,6 +352,54 @@ private static ColumnVector makeExactWidthIntVariantColumn() { field("l", VariantEncoder.int64(1234567890123456789L)))))); } + private static ColumnVector makeArrayVariantColumn() { + VariantEncoder encoder = new VariantEncoder(); + return ColumnVector.fromStructs( + VARIANT_TYPE, + variant(encoder.metadata(), VariantEncoder.array( + VariantEncoder.int8(2), VariantEncoder.int8(1), VariantEncoder.int8(5))), + variant(encoder.metadata(), VariantEncoder.array( + VariantEncoder.int8(9), VariantEncoder.int8(8)))); + } + + private static ColumnVector makeMixedArrayVariantColumn() { + VariantEncoder encoder = new VariantEncoder("a", "b"); + return ColumnVector.fromStructs(VARIANT_TYPE, variant( + encoder.metadata(), + encoder.object(fields(field("a", VariantEncoder.array( + encoder.object(fields(field("b", VariantEncoder.int32(7)))), + encoder.object(fields(field("b", VariantEncoder.int32(42)))))))))); + } + + private static ColumnVector makeFloatVariantColumn() { + VariantEncoder encoder = new VariantEncoder("f", "d"); + return ColumnVector.fromStructs( + VARIANT_TYPE, + variant(encoder.metadata(), encoder.object(fields( + field("f", VariantEncoder.float32(1.25f)), + field("d", VariantEncoder.float64(-2.5))))), + variant(encoder.metadata(), encoder.object(fields( + field("f", VariantEncoder.float32(-4.5f)), + field("d", VariantEncoder.float64(8.125)))))); + } + + private static ColumnVector makeBoolVariantColumn() { + VariantEncoder encoder = new VariantEncoder("b"); + return ColumnVector.fromStructs( + VARIANT_TYPE, + variant(encoder.metadata(), encoder.object(fields( + field("b", VariantEncoder.int8(0))))), + variant(encoder.metadata(), encoder.object(fields( + field("b", VariantEncoder.bool(true))))), + variant(encoder.metadata(), encoder.object(fields( + field("b", VariantEncoder.bool(false))))), + variant(encoder.metadata(), encoder.object(fields( + field("b", VariantEncoder.nullValue())))), + variant(encoder.metadata(), encoder.object(fields( + field("b", VariantEncoder.int8(1))))), + null); + } + @Test void extractStringField() { try (ColumnVector variant = makeXyzVariantColumn(); @@ -389,6 +485,130 @@ void getThenCastFieldValue() { } } + @Test + void extractRootArrayElement() { + try (ColumnVector variant = makeArrayVariantColumn(); + ColumnVector result = VariantUtils.extractVariantField(variant, "$[0]", DType.INT8); + ColumnVector expected = ColumnVector.fromBoxedBytes((byte) 2, (byte) 9)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void getThenCastArrayElement() { + try (ColumnVector variant = makeArrayVariantColumn(); + ColumnVector valueBytes = VariantUtils.getVariantFieldValue(variant, "$[1]"); + ColumnVector result = VariantUtils.castVariantValue(valueBytes, DType.INT8); + ColumnVector expected = ColumnVector.fromBoxedBytes((byte) 1, (byte) 8)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void extractMixedObjectArrayPath() { + try (ColumnVector variant = makeMixedArrayVariantColumn(); + ColumnVector result = VariantUtils.extractVariantField( + variant, "$.a[1].b", DType.INT32); + ColumnVector expected = ColumnVector.fromBoxedInts(42)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void arrayIndexResolutionFailuresProduceNulls() { + try (ColumnVector variant = makeArrayVariantColumn(); + ColumnVector outOfBounds = VariantUtils.extractVariantField( + variant, "$[99]", DType.INT8); + ColumnVector containerMismatch = VariantUtils.extractVariantField( + variant, "$[0][0]", DType.INT8); + ColumnVector expected = ColumnVector.fromBoxedBytes(null, null)) { + assertColumnsAreEqual(expected, outOfBounds); + assertColumnsAreEqual(expected, containerMismatch); + } + } + + @Test + void extractArrayElementFromSlice() { + try (ColumnVector variant = makeArrayVariantColumn(); + ColumnVector sliced = variant.subVector(1, 2); + ColumnVector result = VariantUtils.extractVariantField(sliced, "$[0]", DType.INT8); + ColumnVector expected = ColumnVector.fromBoxedBytes((byte) 9)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void castFloatFields() { + try (ColumnVector variant = makeFloatVariantColumn(); + ColumnVector floatBytes = VariantUtils.getVariantFieldValue(variant, "f"); + ColumnVector doubleBytes = VariantUtils.getVariantFieldValue(variant, "d"); + ColumnVector floats = VariantUtils.castVariantValue(floatBytes, DType.FLOAT32); + ColumnVector doubles = VariantUtils.castVariantValue(doubleBytes, DType.FLOAT64); + ColumnVector expectedFloats = ColumnVector.fromBoxedFloats(1.25f, -4.5f); + ColumnVector expectedDoubles = ColumnVector.fromBoxedDoubles(-2.5, 8.125)) { + assertColumnsAreEqual(expectedFloats, floats); + assertColumnsAreEqual(expectedDoubles, doubles); + } + } + + @Test + void extractFloatFields() { + try (ColumnVector variant = makeFloatVariantColumn(); + ColumnVector floats = VariantUtils.extractVariantField(variant, "f", DType.FLOAT32); + ColumnVector doubles = VariantUtils.extractVariantField(variant, "d", DType.FLOAT64); + ColumnVector expectedFloats = ColumnVector.fromBoxedFloats(1.25f, -4.5f); + ColumnVector expectedDoubles = ColumnVector.fromBoxedDoubles(-2.5, 8.125)) { + assertColumnsAreEqual(expectedFloats, floats); + assertColumnsAreEqual(expectedDoubles, doubles); + } + } + + @Test + void floatWidthMismatchProducesNulls() { + try (ColumnVector variant = makeFloatVariantColumn(); + ColumnVector result = VariantUtils.extractVariantField(variant, "f", DType.FLOAT64); + ColumnVector expected = ColumnVector.fromBoxedDoubles(null, null)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void castFloatNullInputIsPreserved() { + try (ColumnVector values = ColumnVector.fromLists( + BINARY_TYPE, VariantEncoder.float32(1.25f), null); + ColumnVector result = VariantUtils.castVariantValue(values, DType.FLOAT32); + ColumnVector expected = ColumnVector.fromBoxedFloats(1.25f, null)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void castBooleanValues() { + try (ColumnVector values = ColumnVector.fromLists( + BINARY_TYPE, + VariantEncoder.int8(0), + VariantEncoder.bool(true), + VariantEncoder.bool(false), + VariantEncoder.nullValue(), + VariantEncoder.int8(1), + null); + ColumnVector sliced = values.subVector(1, 6); + ColumnVector result = VariantUtils.castVariantValue(sliced, DType.BOOL8); + ColumnVector expected = ColumnVector.fromBoxedBooleans(true, false, null, null, null)) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void extractBooleanFieldFromSlice() { + try (ColumnVector variant = makeBoolVariantColumn(); + ColumnVector sliced = variant.subVector(1, 6); + ColumnVector result = VariantUtils.extractVariantField(sliced, "b", DType.BOOL8); + ColumnVector expected = ColumnVector.fromBoxedBooleans(true, false, null, null, null)) { + assertColumnsAreEqual(expected, result); + } + } + @Test void emptyInputProducesEmptyOutput() { try (ColumnVector variant = ColumnVector.fromStructs(VARIANT_TYPE); @@ -398,6 +618,28 @@ void emptyInputProducesEmptyOutput() { } } + @Test + void emptyInputSupportsFloatOutput() { + try (ColumnVector variant = ColumnVector.fromStructs(VARIANT_TYPE); + ColumnVector result = VariantUtils.extractVariantField(variant, "x", DType.FLOAT32); + ColumnVector expected = ColumnVector.fromBoxedFloats()) { + assertColumnsAreEqual(expected, result); + } + } + + @Test + void emptyInputSupportsBooleanOutput() { + try (ColumnVector values = ColumnVector.fromLists(BINARY_TYPE); + ColumnVector directResult = VariantUtils.castVariantValue(values, DType.BOOL8); + ColumnVector variant = ColumnVector.fromStructs(VARIANT_TYPE); + ColumnVector extractedResult = VariantUtils.extractVariantField( + variant, "x", DType.BOOL8); + ColumnVector expected = ColumnVector.fromBoxedBooleans()) { + assertColumnsAreEqual(expected, directResult); + assertColumnsAreEqual(expected, extractedResult); + } + } + @Test void emptyPathThrows() { try (ColumnVector variant = makeXyzVariantColumn()) { @@ -440,19 +682,48 @@ void unsupportedTargetTypeThrows() { try (ColumnVector variant = makeXyzVariantColumn(); ColumnVector valueBytes = VariantUtils.getVariantFieldValue(variant, "x")) { assertThrows(IllegalArgumentException.class, - () -> VariantUtils.castVariantValue(valueBytes, DType.FLOAT64)); + () -> VariantUtils.castVariantValue(valueBytes, DType.UINT32)); assertThrows(IllegalArgumentException.class, - () -> VariantUtils.extractVariantField(variant, "x", DType.FLOAT64)); + () -> VariantUtils.extractVariantField(variant, "x", DType.UINT32)); } } @Test - void nullCastArgumentsThrow() { - assertThrows(NullPointerException.class, () -> VariantUtils.castVariantValue(null, null)); - assertThrows(NullPointerException.class, - () -> VariantUtils.castVariantValue(null, DType.INT32)); - assertThrows(NullPointerException.class, - () -> VariantUtils.castVariantValue(null, DType.FLOAT64)); + void nullArgumentsThrow() { + try (ColumnVector values = ColumnVector.fromLists(BINARY_TYPE, VariantEncoder.int32(1)); + ColumnVector variant = makeXyzVariantColumn()) { + assertThrows(NullPointerException.class, () -> VariantUtils.castVariantValue(null, null)); + assertThrows(NullPointerException.class, + () -> VariantUtils.castVariantValue(null, DType.INT32)); + assertThrows(NullPointerException.class, + () -> VariantUtils.castVariantValue(null, DType.FLOAT64)); + assertThrows(NullPointerException.class, + () -> VariantUtils.castVariantValue(values, null)); + assertThrows(NullPointerException.class, + () -> VariantUtils.extractVariantField(variant, "x", null)); + } + } + + @Test + void invalidInputShapesThrow() { + try (ColumnVector nonVariant = ColumnVector.fromInts(1, 2, 3)) { + assertThrows(CudfException.class, + () -> VariantUtils.getVariantFieldValue(nonVariant, "x")); + assertThrows(CudfException.class, + () -> VariantUtils.castVariantValue(nonVariant, DType.INT32)); + assertThrows(CudfException.class, + () -> VariantUtils.extractVariantField(nonVariant, "x", DType.INT32)); + } + } + + @Test + void truncatedFloatPayloadProducesNull() { + try (ColumnVector values = ColumnVector.fromLists( + BINARY_TYPE, bytes(VariantEncoder.simple(0x07), 0x00, 0x00)); + ColumnVector result = VariantUtils.castVariantValue(values, DType.FLOAT64); + ColumnVector expected = ColumnVector.fromBoxedDoubles((Double) null)) { + assertColumnsAreEqual(expected, result); + } } @Test