From 30defe1fa670e9a95f1c2d8d2087f9505a9d7ccb Mon Sep 17 00:00:00 2001 From: abdul rawoof Date: Fri, 2 Oct 2026 15:58:22 +0530 Subject: [PATCH] GH-1317: bound FlightData frame field length before allocating --- .../org/apache/arrow/flight/ArrowMessage.java | 22 ++++ .../apache/arrow/flight/TestArrowMessage.java | 101 ++++++++++++++++++ 2 files changed, 123 insertions(+) create mode 100644 flight/flight-core/src/test/java/org/apache/arrow/flight/TestArrowMessage.java 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 marshaller = + ArrowMessage.createMarshaller(allocator); + final Exception e = + assertThrows( + Exception.class, () -> marshaller.parse(new ByteArrayInputStream(frame.toByteArray()))); + assertTrue( + e.getMessage() != null && e.getMessage().contains("exceeds"), + "unexpected failure: " + e.getMessage()); + } + + /** A well-formed field whose length matches the bytes present still parses. */ + @Test + public void frameAcceptsWellFormedField() throws Exception { + final byte[] payload = new byte[] {1, 2, 3, 4}; + final ByteArrayOutputStream frame = new ByteArrayOutputStream(); + writeRawVarint32(frame, APP_METADATA_TAG); + writeRawVarint32(frame, payload.length); + frame.write(payload); + + final MethodDescriptor.Marshaller marshaller = + ArrowMessage.createMarshaller(allocator); + try (ArrowMessage message = marshaller.parse(new ByteArrayInputStream(frame.toByteArray()))) { + assertNotNull(message.getApplicationMetadata()); + } + } + + private static void writeRawVarint32(ByteArrayOutputStream out, int value) { + while (true) { + if ((value & ~0x7F) == 0) { + out.write(value); + return; + } + out.write((value & 0x7F) | 0x80); + value >>>= 7; + } + } +}