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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 26 additions & 6 deletions java/src/main/java/ai/rapids/cudf/VariantUtils.java
Original file line number Diff line number Diff line change
Expand Up @@ -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<DType> 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() {}

Expand All @@ -38,8 +39,11 @@ private static void validateTargetType(DType targetType) {
*
* @param variantStruct Variant materialization: STRUCT(metadata LIST&lt;UINT8&gt;,
* value LIST&lt;UINT8&gt;, 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&lt;UINT8&gt; column of raw encoded Variant values
*/
public static ColumnVector getVariantFieldValue(ColumnView variantStruct, String path) {
Expand All @@ -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");
Expand All @@ -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&lt;UINT8&gt;,
* value LIST&lt;UINT8&gt;, 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) {
Expand Down
287 changes: 279 additions & 8 deletions java/src/test/java/ai/rapids/cudf/VariantUtilsTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,24 @@ private List<Byte> object(Map<String, List<Byte>> fields, List<String> valueOrde
return concat(bytes(header, fields.size()), fieldIds, offsets, encodedValues);
}

@SafeVarargs
private static List<Byte> array(List<Byte>... values) {
List<Byte> offsets = new ArrayList<>(values.length + 1);
List<Byte> encodedValues = new ArrayList<>();
int offset = 0;
for (List<Byte> 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<String> sortedFieldNames(Map<String, List<Byte>> fields) {
List<String> sortedNames = new ArrayList<>(fields.keySet());
sortedNames.sort(VariantEncoder::compareUtf8Unsigned);
Expand Down Expand Up @@ -208,6 +226,36 @@ private static List<Byte> int64(long value) {
(int) ((value >>> 56) & 0xff));
}

private static List<Byte> 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<Byte> 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<Byte> bool(boolean value) {
return bytes(simple(value ? 0x01 : 0x02));
}

private static List<Byte> nullValue() {
return bytes(simple(0x00));
}

private static List<Byte> string(String value) {
byte[] bytes = value.getBytes(StandardCharsets.UTF_8);
if (bytes.length >= (1 << 6)) {
Expand Down Expand Up @@ -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();
Expand Down Expand Up @@ -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);
Expand All @@ -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()) {
Expand Down Expand Up @@ -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
Expand Down
Loading