diff --git a/flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java b/flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java index ab4eab3048..d1c70031ea 100644 --- a/flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java +++ b/flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java @@ -18,12 +18,12 @@ import com.google.common.collect.ImmutableList; import com.google.common.collect.Iterables; -import com.google.common.io.ByteStreams; import com.google.protobuf.ByteString; import com.google.protobuf.CodedInputStream; import com.google.protobuf.CodedOutputStream; import com.google.protobuf.WireFormat; import io.grpc.Drainable; +import io.grpc.KnownLength; import io.grpc.MethodDescriptor.Marshaller; import io.grpc.protobuf.ProtoUtils; import io.netty.buffer.ByteBuf; @@ -281,11 +281,11 @@ public Iterable getBufs() { private static ArrowMessage frame(BufferAllocator allocator, final InputStream stream) { + ArrowBuf body = null; + ArrowBuf appMetadata = null; try { FlightDescriptor descriptor = null; MessageMetadataResult header = null; - ArrowBuf body = null; - ArrowBuf appMetadata = null; while (stream.available() > 0) { final int tagFirstByte = stream.read(); if (tagFirstByte == -1) { @@ -295,25 +295,24 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s switch (tag) { case DESCRIPTOR_TAG: { - int size = readRawVarint32(stream); - byte[] bytes = new byte[size]; - ByteStreams.readFully(stream, bytes); + byte[] bytes = readFieldBytes(stream); descriptor = FlightDescriptor.parseFrom(bytes); break; } case HEADER_TAG: { - int size = readRawVarint32(stream); - byte[] bytes = new byte[size]; - ByteStreams.readFully(stream, bytes); - header = MessageMetadataResult.create(ByteBuffer.wrap(bytes), size); + byte[] bytes = readFieldBytes(stream); + header = MessageMetadataResult.create(ByteBuffer.wrap(bytes), bytes.length); break; } case APP_METADATA_TAG: { - int size = readRawVarint32(stream); - appMetadata = allocator.buffer(size); - GetReadableBuffer.readIntoBuffer(stream, appMetadata, size, ENABLE_ZERO_COPY_READ); + if (appMetadata != null) { + // only read last app metadata. + appMetadata.close(); + appMetadata = null; + } + appMetadata = readFieldBuffer(allocator, stream); break; } case BODY_TAG: @@ -322,9 +321,7 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s body.getReferenceManager().release(); body = null; } - int size = readRawVarint32(stream); - body = allocator.buffer(size); - GetReadableBuffer.readIntoBuffer(stream, body, size, ENABLE_ZERO_COPY_READ); + body = readFieldBuffer(allocator, stream); break; default: @@ -364,6 +361,8 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s } return new ArrowMessage(descriptor, header, appMetadata, body); } catch (Exception ioe) { + // No ArrowMessage takes ownership of the buffers read so far, so release them here. + AutoCloseables.close(ioe, appMetadata, body); throw new RuntimeException(ioe); } } @@ -377,6 +376,71 @@ private static int readRawVarint32(int firstByte, InputStream is) throws IOExcep return CodedInputStream.readRawVarint32(firstByte, is); } + /** + * Read and validate the length prefix of a length-delimited field. + * + *

The length is read straight off the wire, so it must not size an allocation unchecked. It + * can only be compared to the bytes left in the message when the stream is {@link KnownLength}: + * {@code available()} means something else on other streams (the decompressing stream gRPC uses + * for compressed messages reports 1 until EOF), so for those the read itself bounds the + * allocation, see {@link #readFieldBytes}. + */ + private static int readFieldLength(InputStream stream) throws IOException { + final int size = readRawVarint32(stream); + if (size < 0) { + throw new IOException("Malformed FlightData frame: negative field length " + size); + } + if (stream instanceof KnownLength && size > stream.available()) { + throw fieldTooLong(size, stream.available()); + } + return size; + } + + /** + * Read a length-delimited field into a byte array. + * + *

{@link InputStream#readNBytes(int)} grows the array as bytes arrive, so a length prefix + * larger than the message is rejected without allocating the declared length first. + */ + private static byte[] readFieldBytes(InputStream stream) throws IOException { + final int size = readFieldLength(stream); + final byte[] bytes = stream.readNBytes(size); + if (bytes.length != size) { + throw fieldTooLong(size, bytes.length); + } + return bytes; + } + + /** Read a length-delimited field into a new buffer. The caller must release the buffer. */ + private static ArrowBuf readFieldBuffer(BufferAllocator allocator, InputStream stream) + throws IOException { + if (!(stream instanceof KnownLength)) { + // The length can't be checked up front, so only allocate once the bytes have arrived. + final byte[] bytes = readFieldBytes(stream); + final ArrowBuf buf = allocator.buffer(bytes.length); + buf.writeBytes(bytes); + return buf; + } + final int size = readFieldLength(stream); + final ArrowBuf buf = allocator.buffer(size); + try { + GetReadableBuffer.readIntoBuffer(stream, buf, size, ENABLE_ZERO_COPY_READ); + } catch (IOException | RuntimeException e) { + buf.close(); + throw e; + } + return buf; + } + + private static IOException fieldTooLong(int size, int remaining) { + return new IOException( + "Malformed FlightData frame: field length " + + size + + " exceeds " + + remaining + + " bytes remaining in the message"); + } + /** * Convert the ArrowMessage to an InputStream. * diff --git a/flight/flight-core/src/test/java/org/apache/arrow/flight/TestArrowMessage.java b/flight/flight-core/src/test/java/org/apache/arrow/flight/TestArrowMessage.java new file mode 100644 index 0000000000..41e39ad2d6 --- /dev/null +++ b/flight/flight-core/src/test/java/org/apache/arrow/flight/TestArrowMessage.java @@ -0,0 +1,208 @@ +/* + * 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.flight; + +import static org.junit.jupiter.api.Assertions.assertArrayEquals; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import com.google.common.collect.Iterables; +import com.google.protobuf.ByteString; +import com.google.protobuf.CodedOutputStream; +import com.google.protobuf.WireFormat; +import io.grpc.MethodDescriptor; +import io.grpc.internal.ReadableBuffers; +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.IOException; +import java.io.InputStream; +import java.util.Collections; +import java.util.zip.GZIPInputStream; +import java.util.zip.GZIPOutputStream; +import org.apache.arrow.flight.impl.Flight.FlightData; +import org.apache.arrow.memory.ArrowBuf; +import org.apache.arrow.memory.BufferAllocator; +import org.apache.arrow.memory.RootAllocator; +import org.apache.arrow.vector.ipc.message.IpcOption; +import org.apache.arrow.vector.ipc.message.MessageSerializer; +import org.apache.arrow.vector.types.pojo.ArrowType; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.arrow.vector.types.pojo.Schema; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.EnumSource; + +/** Tests for deframing FlightData messages in {@link ArrowMessage}. */ +public class TestArrowMessage { + + private static final int[] LENGTH_DELIMITED_FIELDS = { + FlightData.FLIGHT_DESCRIPTOR_FIELD_NUMBER, + FlightData.DATA_HEADER_FIELD_NUMBER, + FlightData.APP_METADATA_FIELD_NUMBER, + FlightData.DATA_BODY_FIELD_NUMBER + }; + + /** A declared field length far larger than any frame built here. */ + private static final int OVERSIZED_LENGTH = 1 << 20; + + private static final byte[] FIRST = new byte[] {1, 2, 3, 4}; + private static final byte[] LAST = new byte[] {5, 6, 7, 8, 9, 10}; + + /** The kinds of stream gRPC hands to the marshaller. */ + enum StreamType { + /** An uncompressed message: available() is the number of bytes left in the message. */ + KNOWN_LENGTH { + @Override + InputStream open(byte[] frame) { + return ReadableBuffers.openStream(ReadableBuffers.wrap(frame), true); + } + }, + /** A compressed message: available() is 1 until EOF, however long the message is. */ + COMPRESSED { + @Override + InputStream open(byte[] frame) throws IOException { + final ByteArrayOutputStream compressed = new ByteArrayOutputStream(); + try (GZIPOutputStream gzip = new GZIPOutputStream(compressed)) { + gzip.write(frame); + } + return new GZIPInputStream(new ByteArrayInputStream(compressed.toByteArray())); + } + }; + + abstract InputStream open(byte[] frame) throws IOException; + } + + private BufferAllocator allocator; + private MethodDescriptor.Marshaller marshaller; + + @BeforeEach + public void setUp() { + allocator = new RootAllocator(Long.MAX_VALUE); + marshaller = ArrowMessage.createMarshaller(allocator); + } + + @AfterEach + public void tearDown() { + // Fails if a test leaked a buffer. + allocator.close(); + } + + /** A well-formed frame parses whichever kind of stream it arrives on. */ + @ParameterizedTest + @EnumSource + public void frameAcceptsWellFormedFrame(StreamType streamType) throws Exception { + final Schema schema = + new Schema(Collections.singletonList(Field.nullable("foo", new ArrowType.Int(32, true)))); + final FlightData data = + FlightData.newBuilder() + .setFlightDescriptor(FlightDescriptor.command(FIRST).toProtocol()) + .setDataHeader( + ByteString.copyFrom(MessageSerializer.serializeMetadata(schema, IpcOption.DEFAULT))) + .setAppMetadata(ByteString.copyFrom(LAST)) + .build(); + + try (ArrowMessage message = marshaller.parse(streamType.open(data.toByteArray()))) { + assertEquals(data.getFlightDescriptor(), message.getDescriptor()); + assertEquals(schema, message.asSchema()); + assertArrayEquals(LAST, toByteArray(message.getApplicationMetadata())); + } + } + + /** A repeated buffer field keeps the last occurrence and releases the earlier one. */ + @ParameterizedTest + @EnumSource + public void frameKeepsLastOfRepeatedField(StreamType streamType) throws Exception { + final ByteArrayOutputStream frame = new ByteArrayOutputStream(); + for (byte[] value : new byte[][] {FIRST, LAST}) { + FlightData.newBuilder() + .setAppMetadata(ByteString.copyFrom(value)) + .setDataBody(ByteString.copyFrom(value)) + .build() + .writeTo(frame); + } + + try (ArrowMessage message = marshaller.parse(streamType.open(frame.toByteArray()))) { + assertArrayEquals(LAST, toByteArray(message.getApplicationMetadata())); + assertArrayEquals(LAST, toByteArray(Iterables.getOnlyElement(message.getBufs()))); + } + assertEquals(0, allocator.getAllocatedMemory()); + } + + /** + * A field declaring more bytes than the frame holds is rejected, without first allocating a + * buffer of the declared length. + */ + @ParameterizedTest + @EnumSource + public void frameRejectsOversizedFieldLength(StreamType streamType) throws Exception { + for (int field : LENGTH_DELIMITED_FIELDS) { + assertRejected(streamType, fieldPrefix(field, OVERSIZED_LENGTH)); + } + assertEquals(0, allocator.getPeakMemoryAllocation()); + } + + /** A negative length, which a 5-byte varint can encode, is rejected. */ + @ParameterizedTest + @EnumSource + public void frameRejectsNegativeFieldLength(StreamType streamType) throws Exception { + for (int field : LENGTH_DELIMITED_FIELDS) { + assertRejected(streamType, fieldPrefix(field, -1)); + } + } + + /** Buffers read for earlier fields are released when a later field is rejected. */ + @ParameterizedTest + @EnumSource + public void frameReleasesBuffersWhenLaterFieldIsRejected(StreamType streamType) throws Exception { + for (int field : LENGTH_DELIMITED_FIELDS) { + final ByteArrayOutputStream frame = new ByteArrayOutputStream(); + FlightData.newBuilder() + .setAppMetadata(ByteString.copyFrom(FIRST)) + .setDataBody(ByteString.copyFrom(LAST)) + .build() + .writeTo(frame); + frame.write(fieldPrefix(field, OVERSIZED_LENGTH)); + + assertRejected(streamType, frame.toByteArray()); + assertEquals(0, allocator.getAllocatedMemory()); + } + } + + private void assertRejected(StreamType streamType, byte[] frame) throws IOException { + final InputStream stream = streamType.open(frame); + final RuntimeException e = assertThrows(RuntimeException.class, () -> marshaller.parse(stream)); + assertInstanceOf(IOException.class, e.getCause()); + } + + /** The tag and length prefix of a length-delimited field, without any content. */ + private static byte[] fieldPrefix(int fieldNumber, int length) throws IOException { + final ByteArrayOutputStream frame = new ByteArrayOutputStream(); + final CodedOutputStream out = CodedOutputStream.newInstance(frame); + out.writeTag(fieldNumber, WireFormat.WIRETYPE_LENGTH_DELIMITED); + out.writeUInt32NoTag(length); + out.flush(); + return frame.toByteArray(); + } + + private static byte[] toByteArray(ArrowBuf buf) { + final byte[] bytes = new byte[(int) buf.readableBytes()]; + buf.getBytes(0, bytes); + return bytes; + } +}