diff --git a/vector/src/main/java/org/apache/arrow/vector/VectorLoader.java b/vector/src/main/java/org/apache/arrow/vector/VectorLoader.java index 9b9a890346..9e2992761f 100644 --- a/vector/src/main/java/org/apache/arrow/vector/VectorLoader.java +++ b/vector/src/main/java/org/apache/arrow/vector/VectorLoader.java @@ -106,11 +106,15 @@ private void loadBuffers( Iterator nodes, CompressionCodec codec, Iterator variadicBufferCounts) { + FieldVector storageVector = vector; + while (storageVector instanceof ExtensionTypeVector) { + storageVector = ((ExtensionTypeVector) storageVector).getUnderlyingVector(); + } checkArgument(nodes.hasNext(), "no more field nodes for field %s and vector %s", field, vector); ArrowFieldNode fieldNode = nodes.next(); - // variadicBufferLayoutCount will be 0 for vectors of a type except BaseVariableWidthViewVector + // Only view storage has variadic buffers. long variadicBufferLayoutCount = 0; - if (vector instanceof BaseVariableWidthViewVector) { + if (storageVector instanceof BaseVariableWidthViewVector) { if (variadicBufferCounts.hasNext()) { variadicBufferLayoutCount = variadicBufferCounts.next(); } else { diff --git a/vector/src/main/java/org/apache/arrow/vector/VectorUnloader.java b/vector/src/main/java/org/apache/arrow/vector/VectorUnloader.java index 342f210b82..66d0debff0 100644 --- a/vector/src/main/java/org/apache/arrow/vector/VectorUnloader.java +++ b/vector/src/main/java/org/apache/arrow/vector/VectorUnloader.java @@ -104,14 +104,18 @@ private void appendNodes( List nodes, List buffers, List variadicBufferCounts) { + FieldVector storageVector = vector; + while (storageVector instanceof ExtensionTypeVector) { + storageVector = ((ExtensionTypeVector) storageVector).getUnderlyingVector(); + } nodes.add( new ArrowFieldNode(vector.getValueCount(), includeNullCount ? vector.getNullCount() : -1)); List fieldBuffers = vector.getFieldBuffers(); - long variadicBufferCount = getVariadicBufferCount(vector); + long variadicBufferCount = getVariadicBufferCount(storageVector); int expectedBufferCount = (int) (TypeLayout.getTypeBufferCount(vector.getField().getType()) + variadicBufferCount); // only update variadicBufferCounts for vectors that have variadic buffers - if (vector instanceof BaseVariableWidthViewVector) { + if (storageVector instanceof BaseVariableWidthViewVector) { variadicBufferCounts.add(variadicBufferCount); } if (fieldBuffers.size() != expectedBufferCount) { diff --git a/vector/src/main/java/org/apache/arrow/vector/extension/JsonType.java b/vector/src/main/java/org/apache/arrow/vector/extension/JsonType.java new file mode 100644 index 0000000000..be3ebb4f7b --- /dev/null +++ b/vector/src/main/java/org/apache/arrow/vector/extension/JsonType.java @@ -0,0 +1,132 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.arrow.vector.extension; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.databind.DeserializationFeature; +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import java.util.Collections; +import java.util.Objects; +import org.apache.arrow.memory.BufferAllocator; +import org.apache.arrow.vector.FieldVector; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.ExtensionTypeRegistry; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.types.pojo.FieldType; + +/** + * Canonical extension type for UTF-8 encoded RFC 8259 JSON values. + * + *

The storage type is {@link ArrowType.Utf8}, {@link ArrowType.LargeUtf8}, or {@link + * ArrowType.Utf8View}. Values use the corresponding string vector; this type does not parse or + * validate individual JSON values. + * + *

Register the type before reading schemas containing {@code arrow.json}: + * + *

{@code
+ * JsonType.ensureRegistered();
+ * Field field = Field.nullable("json", new JsonType(ArrowType.Utf8.INSTANCE));
+ * try (JsonVector vector = (JsonVector) field.createVector(allocator)) {
+ *   VarCharVector storage = (VarCharVector) vector.getUnderlyingVector();
+ *   storage.setSafe(0, "{}".getBytes(java.nio.charset.StandardCharsets.UTF_8));
+ *   vector.setValueCount(1);
+ *   Text value = vector.getObject(0);
+ * }
+ * }
+ */ +public class JsonType extends ArrowType.ExtensionType { + public static final String EXTENSION_NAME = "arrow.json"; + private static final ObjectMapper MAPPER = + new ObjectMapper().enable(DeserializationFeature.FAIL_ON_TRAILING_TOKENS); + private final ArrowType storageType; + + /** Register a prototype that can deserialize all supported JSON storage types. */ + public static void ensureRegistered() { + ExtensionTypeRegistry.register(new JsonType(ArrowType.Utf8.INSTANCE)); + } + + /** + * Create a JSON type backed by the specified string type. + * + * @param storageType Utf8, LargeUtf8, or Utf8View + * @throws IllegalArgumentException if the storage type is not a supported string type + */ + public JsonType(ArrowType storageType) { + Objects.requireNonNull(storageType, "storageType"); + if (!(storageType instanceof ArrowType.Utf8) + && !(storageType instanceof ArrowType.LargeUtf8) + && !(storageType instanceof ArrowType.Utf8View)) { + throw new IllegalArgumentException( + "arrow.json requires Utf8, LargeUtf8, or Utf8View storage, got " + storageType); + } + this.storageType = storageType; + } + + @Override + public ArrowType storageType() { + return storageType; + } + + @Override + public String extensionName() { + return EXTENSION_NAME; + } + + @Override + public boolean extensionEquals(ExtensionType other) { + return other instanceof JsonType && storageType.equals(other.storageType()); + } + + @Override + public String serialize() { + return ""; + } + + @Override + public ArrowType deserialize(ArrowType storageType, String serializedData) { + JsonType type = new JsonType(storageType); + if (serializedData == null) { + throw new InvalidExtensionMetadataException("arrow.json metadata must not be null"); + } + if (!serializedData.isEmpty()) { + try { + JsonNode metadata = MAPPER.readTree(serializedData); + if (metadata == null || !metadata.isObject()) { + throw new InvalidExtensionMetadataException("arrow.json metadata must be a JSON object"); + } + } catch (JsonProcessingException e) { + throw new InvalidExtensionMetadataException("arrow.json metadata is invalid", e); + } + } + return type; + } + + @Override + public boolean isComplex() { + return false; + } + + @Override + public FieldVector getNewVector(String name, FieldType fieldType, BufferAllocator allocator) { + Field field = new Field(name, fieldType, Collections.emptyList()); + FieldType storageFieldType = + new FieldType(fieldType.isNullable(), storageType, fieldType.getDictionary(), null); + FieldVector storage = storageFieldType.createNewSingleVector(name, allocator, null); + return new JsonVector(field, allocator, storage); + } +} diff --git a/vector/src/main/java/org/apache/arrow/vector/extension/JsonVector.java b/vector/src/main/java/org/apache/arrow/vector/extension/JsonVector.java new file mode 100644 index 0000000000..76c8579d06 --- /dev/null +++ b/vector/src/main/java/org/apache/arrow/vector/extension/JsonVector.java @@ -0,0 +1,126 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.arrow.vector.extension; + +import org.apache.arrow.memory.BufferAllocator; +import org.apache.arrow.memory.util.hash.ArrowBufHasher; +import org.apache.arrow.vector.ExtensionTypeVector; +import org.apache.arrow.vector.FieldVector; +import org.apache.arrow.vector.ValueIterableVector; +import org.apache.arrow.vector.ValueVector; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.util.CallBack; +import org.apache.arrow.vector.util.Text; +import org.apache.arrow.vector.util.TransferPair; + +/** + * A JSON extension vector backed by a string vector. + * + *

Use {@link Field#createVector(BufferAllocator)} with a {@link JsonType} field to create an + * instance. Write UTF-8 JSON through {@link #getUnderlyingVector()}; values are not parsed or + * validated. + */ +public class JsonVector extends ExtensionTypeVector + implements ValueIterableVector { + private final Field field; + + JsonVector(Field field, BufferAllocator allocator, FieldVector underlyingVector) { + super(field, allocator, underlyingVector); + this.field = field; + } + + @Override + public Field getField() { + return field; + } + + @Override + public Text getObject(int index) { + return (Text) getUnderlyingVector().getObject(index); + } + + @Override + public TransferPair getTransferPair(BufferAllocator allocator) { + return getTransferPair(field, allocator); + } + + @Override + public TransferPair getTransferPair(String name, BufferAllocator allocator) { + return getTransferPair(new Field(name, field.getFieldType(), field.getChildren()), allocator); + } + + @Override + public TransferPair getTransferPair(String name, BufferAllocator allocator, CallBack callBack) { + return getTransferPair(name, allocator); + } + + @Override + public TransferPair getTransferPair(Field targetField, BufferAllocator allocator) { + return makeTransferPair(targetField.createVector(allocator)); + } + + @Override + public TransferPair getTransferPair( + Field targetField, BufferAllocator allocator, CallBack callBack) { + return getTransferPair(targetField, allocator); + } + + @Override + public TransferPair makeTransferPair(ValueVector target) { + return new TransferImpl((JsonVector) target); + } + + @Override + public int hashCode(int index) { + return hashCode(index, null); + } + + @Override + public int hashCode(int index, ArrowBufHasher hasher) { + return getUnderlyingVector().hashCode(index, hasher); + } + + private class TransferImpl implements TransferPair { + private final JsonVector to; + private final TransferPair storagePair; + + TransferImpl(JsonVector to) { + this.to = to; + this.storagePair = getUnderlyingVector().makeTransferPair(to.getUnderlyingVector()); + } + + @Override + public void transfer() { + storagePair.transfer(); + } + + @Override + public void splitAndTransfer(int startIndex, int length) { + storagePair.splitAndTransfer(startIndex, length); + } + + @Override + public JsonVector getTo() { + return to; + } + + @Override + public void copyValueSafe(int fromIndex, int toIndex) { + storagePair.copyValueSafe(fromIndex, toIndex); + } + } +} diff --git a/vector/src/test/java/org/apache/arrow/vector/TestJsonType.java b/vector/src/test/java/org/apache/arrow/vector/TestJsonType.java new file mode 100644 index 0000000000..ebad0d148f --- /dev/null +++ b/vector/src/test/java/org/apache/arrow/vector/TestJsonType.java @@ -0,0 +1,260 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.arrow.vector; + +import static org.apache.arrow.vector.testing.ValueVectorDataPopulator.setVector; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.nio.ByteBuffer; +import java.nio.charset.StandardCharsets; +import java.util.Collections; +import java.util.stream.Stream; +import org.apache.arrow.memory.BufferAllocator; +import org.apache.arrow.memory.RootAllocator; +import org.apache.arrow.vector.extension.InvalidExtensionMetadataException; +import org.apache.arrow.vector.extension.JsonType; +import org.apache.arrow.vector.extension.JsonVector; +import org.apache.arrow.vector.ipc.ArrowStreamReader; +import org.apache.arrow.vector.ipc.ArrowStreamWriter; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.ExtensionTypeRegistry; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.types.pojo.FieldType; +import org.apache.arrow.vector.types.pojo.Schema; +import org.apache.arrow.vector.util.Text; +import org.apache.arrow.vector.util.TransferPair; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.MethodSource; +import org.junit.jupiter.params.provider.NullSource; +import org.junit.jupiter.params.provider.ValueSource; + +class TestJsonType { + BufferAllocator allocator; + ArrowType.ExtensionType previousType; + + @BeforeEach + void beforeEach() { + allocator = new RootAllocator(); + previousType = ExtensionTypeRegistry.lookup(JsonType.EXTENSION_NAME); + JsonType.ensureRegistered(); + } + + @AfterEach + void afterEach() { + ExtensionTypeRegistry.unregister(new JsonType(ArrowType.Utf8.INSTANCE)); + if (previousType != null) { + ExtensionTypeRegistry.register(previousType); + } + allocator.close(); + } + + static Stream storageTypes() { + return Stream.of( + ArrowType.Utf8.INSTANCE, ArrowType.LargeUtf8.INSTANCE, ArrowType.Utf8View.INSTANCE); + } + + @ParameterizedTest + @MethodSource("storageTypes") + void testRoundTrip(ArrowType storage) { + JsonType type = new JsonType(storage); + assertEquals("arrow.json", type.extensionName()); + assertEquals(storage, type.storageType()); + assertFalse(type.isComplex()); + assertEquals("", type.serialize()); + assertEquals(type, type.deserialize(storage, type.serialize())); + assertNotEquals( + type, + new JsonType( + storage instanceof ArrowType.Utf8 + ? ArrowType.LargeUtf8.INSTANCE + : ArrowType.Utf8.INSTANCE)); + } + + @ParameterizedTest + @ValueSource(strings = {"", "{}", " { } ", "{\"future\": 1}"}) + void testDeserializeValid(String metadata) { + JsonType type = new JsonType(ArrowType.Utf8.INSTANCE); + assertEquals(type, type.deserialize(type.storageType(), metadata)); + } + + @ParameterizedTest + @NullSource + @ValueSource(strings = {" ", "null", "[]", "1", "true", "\"json\"", "{", "{} {}", "{} trailing"}) + void testInvalidMetadata(String metadata) { + assertThrows( + InvalidExtensionMetadataException.class, + () -> new JsonType(ArrowType.Utf8.INSTANCE).deserialize(ArrowType.Utf8.INSTANCE, metadata)); + } + + @Test + void testInvalidStorage() { + for (ArrowType storage : + new ArrowType[] { + ArrowType.Binary.INSTANCE, + ArrowType.LargeBinary.INSTANCE, + ArrowType.BinaryView.INSTANCE, + ArrowType.Null.INSTANCE, + new ArrowType.Int(32, true) + }) { + assertThrows(IllegalArgumentException.class, () -> new JsonType(storage)); + assertThrows( + IllegalArgumentException.class, + () -> new JsonType(ArrowType.Utf8.INSTANCE).deserialize(storage, "")); + } + } + + @ParameterizedTest + @MethodSource("storageTypes") + void testSchemaRoundTrip(ArrowType storage) { + for (boolean nullable : new boolean[] {false, true}) { + Field field = field(storage, nullable); + Schema schema = new Schema(Collections.singletonList(field)); + assertEquals(schema, Schema.deserializeMessage(ByteBuffer.wrap(schema.serializeAsMessage()))); + } + } + + // Generated with PyArrow 24.0.0 using pa.schema([pa.field("json", pa.json_(t), + // nullable=False, metadata={"custom": "preserved"}) for t in + // [pa.string(), pa.large_string(), pa.string_view()]]).serialize(). + @Test + void testPyArrowSchema() throws IOException { + Schema schema; + try (InputStream input = getClass().getResourceAsStream("/pyarrow_json_schema.arrow")) { + schema = Schema.deserializeMessage(ByteBuffer.wrap(input.readAllBytes())); + } + ArrowType[] storage = storageTypes().toArray(ArrowType[]::new); + for (int i = 0; i < storage.length; i++) { + assertEquals(field(storage[i], false), schema.getFields().get(i)); + } + } + + private static Field field(ArrowType storage, boolean nullable) { + return new Field( + "json", + new FieldType( + nullable, new JsonType(storage), null, Collections.singletonMap("custom", "preserved")), + Collections.emptyList()); + } + + @ParameterizedTest + @MethodSource("storageTypes") + void testTransfer(ArrowType storage) { + Field field = field(storage, true); + try (JsonVector source = (JsonVector) field.createVector(allocator)) { + byte[] bytes = + "{\"key\":\"value longer than twelve bytes\"}".getBytes(StandardCharsets.UTF_8); + setVector((VariableWidthFieldVector) source.getUnderlyingVector(), null, bytes, null); + TransferPair split = source.getTransferPair("copy", allocator); + try (JsonVector target = assertInstanceOf(JsonVector.class, split.getTo())) { + split.splitAndTransfer(1, 2); + assertEquals("copy", target.getName()); + assertEquals(field.getFieldType(), target.getField().getFieldType()); + assertEquals(new Text(bytes), target.getObject(0)); + assertTrue(target.isNull(1)); + } + try (JsonVector target = (JsonVector) field.createVector(allocator)) { + TransferPair copy = source.makeTransferPair(target); + copy.copyValueSafe(1, 0); + copy.copyValueSafe(2, 1); + target.setValueCount(2); + assertEquals(new Text(bytes), target.getObject(0)); + assertTrue(target.isNull(1)); + } + TransferPair transfer = source.getTransferPair(allocator); + try (JsonVector target = assertInstanceOf(JsonVector.class, transfer.getTo())) { + transfer.transfer(); + assertEquals(field, target.getField()); + assertTrue(target.isNull(0)); + assertEquals(new Text(bytes), target.getObject(1)); + assertTrue(target.isNull(2)); + assertEquals(3, target.getValueCount()); + assertEquals(0, source.getValueCount()); + } + } + } + + @ParameterizedTest + @MethodSource("storageTypes") + void testVectorIpcRoundTrip(ArrowType storage) throws IOException { + Field field = field(storage, true); + byte[] serialized = writeStream(field); + try (ArrowStreamReader reader = + new ArrowStreamReader(new ByteArrayInputStream(serialized), allocator)) { + assertTrue(reader.loadNextBatch()); + JsonVector vector = + assertInstanceOf(JsonVector.class, reader.getVectorSchemaRoot().getVector(0)); + assertEquals(field, vector.getField()); + assertValues(vector); + } + } + + @ParameterizedTest + @MethodSource("storageTypes") + void testReadUnderlyingType(ArrowType storage) throws IOException { + Field field = field(storage, true); + byte[] serialized = writeStream(field); + ExtensionTypeRegistry.unregister((JsonType) field.getType()); + try (ArrowStreamReader reader = + new ArrowStreamReader(new ByteArrayInputStream(serialized), allocator)) { + assertTrue(reader.loadNextBatch()); + FieldVector vector = reader.getVectorSchemaRoot().getVector(0); + assertEquals(storage, vector.getField().getType()); + assertEquals(field.getMetadata(), vector.getField().getMetadata()); + assertValues(vector); + } + } + + private static final String JSON = "{\"message\":\"你好, a JSON value longer than twelve bytes\"}"; + + private byte[] writeStream(Field field) throws IOException { + ByteArrayOutputStream out = new ByteArrayOutputStream(); + try (VectorSchemaRoot root = + VectorSchemaRoot.create(new Schema(Collections.singletonList(field)), allocator); + ArrowStreamWriter writer = new ArrowStreamWriter(root, null, out)) { + JsonVector vector = (JsonVector) root.getVector(0); + setVector( + (VariableWidthFieldVector) vector.getUnderlyingVector(), + JSON.getBytes(StandardCharsets.UTF_8), + null, + "null".getBytes(StandardCharsets.UTF_8)); + root.setRowCount(3); + writer.start(); + writer.writeBatch(); + writer.end(); + } + return out.toByteArray(); + } + + private static void assertValues(FieldVector vector) { + assertEquals(3, vector.getValueCount()); + assertEquals(new Text(JSON), vector.getObject(0)); + assertTrue(vector.isNull(1)); + assertEquals(new Text("null"), vector.getObject(2)); + } +} diff --git a/vector/src/test/resources/pyarrow_json_schema.arrow b/vector/src/test/resources/pyarrow_json_schema.arrow new file mode 100644 index 0000000000..6c44bff645 Binary files /dev/null and b/vector/src/test/resources/pyarrow_json_schema.arrow differ