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..9aae8e78e8 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 @@ -296,6 +296,7 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s case DESCRIPTOR_TAG: { int size = readRawVarint32(stream); + checkFieldLength(size, stream); byte[] bytes = new byte[size]; ByteStreams.readFully(stream, bytes); descriptor = FlightDescriptor.parseFrom(bytes); @@ -304,6 +305,7 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s case HEADER_TAG: { int size = readRawVarint32(stream); + checkFieldLength(size, stream); byte[] bytes = new byte[size]; ByteStreams.readFully(stream, bytes); header = MessageMetadataResult.create(ByteBuffer.wrap(bytes), size); @@ -312,6 +314,7 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s case APP_METADATA_TAG: { int size = readRawVarint32(stream); + checkFieldLength(size, stream); appMetadata = allocator.buffer(size); GetReadableBuffer.readIntoBuffer(stream, appMetadata, size, ENABLE_ZERO_COPY_READ); break; @@ -323,6 +326,7 @@ private static ArrowMessage frame(BufferAllocator allocator, final InputStream s body = null; } int size = readRawVarint32(stream); + checkFieldLength(size, stream); body = allocator.buffer(size); GetReadableBuffer.readIntoBuffer(stream, body, size, ENABLE_ZERO_COPY_READ); break; @@ -377,6 +381,24 @@ private static int readRawVarint32(int firstByte, InputStream is) throws IOExcep return CodedInputStream.readRawVarint32(firstByte, is); } + /** + * Reject a field whose declared length is negative or larger than the bytes left in the message. + * + *
The length prefix is read straight off the wire, and a field can never be longer than the
+ * bytes still buffered for the message. Without this check an oversized value drives an unbounded
+ * allocation before any content is read; the {@code new byte[size]} paths above do so on the JVM
+ * heap, bypassing the {@link BufferAllocator} limit entirely.
+ */
+ private static void checkFieldLength(int size, InputStream stream) throws IOException {
+ final int remaining = stream.available();
+ if (size < 0 || size > remaining) {
+ throw new IOException(
+ String.format(
+ "Malformed FlightData frame: field length %d exceeds %d bytes remaining in the message",
+ size, remaining));
+ }
+ }
+
/**
* 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..151f7eefa2
--- /dev/null
+++ b/flight/flight-core/src/test/java/org/apache/arrow/flight/TestArrowMessage.java
@@ -0,0 +1,101 @@
+/*
+ * 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.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+
+import com.google.protobuf.WireFormat;
+import io.grpc.MethodDescriptor;
+import java.io.ByteArrayInputStream;
+import java.io.ByteArrayOutputStream;
+import org.apache.arrow.flight.impl.Flight.FlightData;
+import org.apache.arrow.memory.BufferAllocator;
+import org.apache.arrow.memory.RootAllocator;
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.BeforeEach;
+import org.junit.jupiter.api.Test;
+
+public class TestArrowMessage {
+
+ private static final int HEADER_TAG =
+ (FlightData.DATA_HEADER_FIELD_NUMBER << 3) | WireFormat.WIRETYPE_LENGTH_DELIMITED;
+ private static final int APP_METADATA_TAG =
+ (FlightData.APP_METADATA_FIELD_NUMBER << 3) | WireFormat.WIRETYPE_LENGTH_DELIMITED;
+
+ private BufferAllocator allocator;
+
+ @BeforeEach
+ public void setUp() {
+ allocator = new RootAllocator(Long.MAX_VALUE);
+ }
+
+ @AfterEach
+ public void tearDown() {
+ allocator.close();
+ }
+
+ /**
+ * A field whose declared length is far larger than the bytes actually present in the frame must
+ * be rejected before anything is allocated for it, rather than driving an allocation sized by the
+ * attacker-controlled length prefix.
+ */
+ @Test
+ public void frameRejectsOversizedFieldLength() {
+ final ByteArrayOutputStream frame = new ByteArrayOutputStream();
+ writeRawVarint32(frame, HEADER_TAG);
+ // Claim a much larger length than the (zero) bytes that follow.
+ writeRawVarint32(frame, 1 << 20);
+
+ final MethodDescriptor.Marshaller