jbonofre commented on code in PR #1318: URL: https://github.com/apache/arrow-java/pull/1318#discussion_r4171477314
########## 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() { Review Comment: Both tests go through `ByteArrayInputStream`, whose `available()` is exact, so they can't see the compressed-stream case. Could you add: - a well-formed frame wrapped in a `GZIPInputStream` (this fails with the current check) - a reject case with a valid `app_metadata` field before the oversized one, so `allocator.close()` in `tearDown` catches a leaked buffer. As written nothing is allocated before the bad field, so "rejected before anything is allocated" isn't actually asserted - a negative length (5-byte varint), and at least one of `DESCRIPTOR` / `BODY`. ########## 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); Review Comment: Nit: the `readRawVarint32` + `checkFieldLength` pair is repeated at four call sites. A single `readFieldLength(stream)` helper that reads and validates would make it impossible to add a new length-delimited case and forget the check. ########## 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<ArrowMessage> 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()); Review Comment: This accepts any `Exception` who message contains "exceeds". It passes because `frame()` wraps the `IOException` in a `RuntimeException`, so `getMessage()` is really the cause's `toString()`. Asserting on the type is sturdier: ```java final RuntimeException e = assertThrows( RuntimeException.class, () -> marshaller.parse(new ByteArrayInputStream(frame.toByteArray()))); assertInstanceOf(IOException.class, e.getCause()); ``` ########## flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java: ########## @@ -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); Review Comment: This buffer (and `body` below) is leaked when a later field is rejected: `checkFieldLength` throws, and the `catch` at the end of `frame()` only wraps and rethrows. A frame of `[app_metadata, N valid bytes][data_header, len=0x7FFFFFFF]` leaks N bytes of direct memory per message, which a peer can repeat until the allocator is exhausted. That is the same DoS this PR is addressing. The EOF path had this leak before, but since we are adding an explicit reject path it should release `appMetadata` and `body` on failure. Related and pre-existing: a repeated `app_metadata` field overwrites the previous buffer without releasing it, unlike the `BODY` case just below. Could you release the earlier one here too? ########## flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java: ########## @@ -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. + * + * <p>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)); Review Comment: Nit: two small things here: - for `size < 0` it reads "field length -1 exceeds 10 bytes remaining", which is misleading when diagnosing a corrupt length prefix. A separate message for the negative case would help. - `String.format` without a `Locale` renders `%d` with the default locale's digits. #1311 just moved the JDBC interval formatting `Locale.ROOT` for this reason. `String.format(Locale.ROOT, ...)` or plain concatenation would be consistent with that. ########## flight/flight-core/src/main/java/org/apache/arrow/flight/ArrowMessage.java: ########## @@ -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. + * + * <p>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(); Review Comment: `available()` is not the number of bytes left in the message in general. It is exact for gRPC's uncompressed `BufferInputStream`, but when the peer compresses messages gRPC hands `parse()` a decompressing stream, and `InflaterInputStream.available()` returns 1 until EOF and 0 after. With this check, every compressed FlightData field longer than 1 byte is rejected as malformed, so DoGet/DoPut/DoExchange fail as soon as a peer uses e.g. `withCompression("gzip")`. That works today: #742 added the `tagFirstByte == -1` guard a few lines up for exactly this stream type. Could we apply the `available()` bound only when the stream is `io.grpc.KnownLength`? ```java if (size < 0) { throw new IOException("Malformed FlightData frame: negative field length " + size); } if (stream instanceof KnownLength && size > stream.available()) { throw new IOException(...); } ``` For other steams the length can't be checked up front, so the read itself needs to be incremental, e.g. `stream.readBytes(size)` plus a length check instead of `new bytes[size]` + `readFully`, and allocating the `ArrowBuf` once the bytes have arrived. That way the allocation follows the bytes actually received for every stream type. ########## 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<ArrowMessage> 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<ArrowMessage> marshaller = + ArrowMessage.createMarshaller(allocator); + try (ArrowMessage message = marshaller.parse(new ByteArrayInputStream(frame.toByteArray()))) { + assertNotNull(message.getApplicationMetadata()); + } + } + + private static void writeRawVarint32(ByteArrayOutputStream out, int value) { Review Comment: Nit: protobuf already provides this. `CodedOutputStream.writeTag(FlightData.DATA_HEADER_FIELD_NUMBER, WireFormat.WIRETYPE_LENGTH_DELIMITED)` and `writeUInt32NoTag(...)` cover the malformed frame, which also removes the need to re-declare `HEADER_TAG` / `APP_METADATA_TAG`. The well-formed case can be `FlightData.newBuilder().setAppMetadata(...).build().toByteArray()`, as the marshaller tests in `TestBasicOperation` do. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
