From a2122da02d615b515d1481ce8bed9822c31670e1 Mon Sep 17 00:00:00 2001 From: Akira Ajisaka Date: Fri, 18 Sep 2026 23:48:16 +0900 Subject: [PATCH] Fix integer overflow in GcmTransportCipher --- .../network/crypto/GcmTransportCipher.java | 19 +++++-- .../network/crypto/GcmAuthEngineSuite.java | 55 +++++++++++++++++++ 2 files changed, 68 insertions(+), 6 deletions(-) diff --git a/common/network-common/src/main/java/org/apache/spark/network/crypto/GcmTransportCipher.java b/common/network-common/src/main/java/org/apache/spark/network/crypto/GcmTransportCipher.java index 654bfde125871..e1a03186ea3b5 100644 --- a/common/network-common/src/main/java/org/apache/spark/network/crypto/GcmTransportCipher.java +++ b/common/network-common/src/main/java/org/apache/spark/network/crypto/GcmTransportCipher.java @@ -39,7 +39,8 @@ public class GcmTransportCipher implements TransportCipher { private static final String HKDF_ALG = "HmacSha256"; - private static final int LENGTH_HEADER_BYTES = 8; + @VisibleForTesting + static final int LENGTH_HEADER_BYTES = 8; @VisibleForTesting static final int CIPHERTEXT_BUFFER_SIZE = 32 * 1024; // 32KB // Maximum plaintext bytes to accumulate before flushing to downstream handlers, even @@ -212,7 +213,7 @@ public boolean release(int decrement) { @Override public long transferTo(WritableByteChannel target, long position) throws IOException { - int transferredThisCall = 0; + long transferredThisCall = 0; // If the header has is not empty, try to write it out to the target. if (headerByteBuffer.hasRemaining()) { int written = target.write(headerByteBuffer); @@ -369,8 +370,9 @@ private boolean initializeExpectedLength(ByteBuf ciphertextNettyBuf) { } expectedLengthBuffer.flip(); expectedLength = expectedLengthBuffer.getLong(); - if (expectedLength < 0) { - throw new IllegalStateException("Invalid expected ciphertext length."); + if (expectedLength < LENGTH_HEADER_BYTES + (long) headerLength) { + throw new IllegalStateException( + "Invalid expected ciphertext length: " + expectedLength); } ciphertextRead += LENGTH_HEADER_BYTES; } @@ -442,8 +444,13 @@ public void channelRead(ChannelHandlerContext ctx, Object ciphertextMessage) int readableBytes = Math.min( nettyBufReadableBytes, ciphertextBuffer.remaining()); - int expectedRemaining = (int) (expectedLength - ciphertextRead); - int bytesToRead = Math.min(readableBytes, expectedRemaining); + long expectedRemaining = expectedLength - ciphertextRead; + if (expectedRemaining <= 0) { + throw new IllegalStateException( + "Invalid ciphertext state: expectedLength=" + expectedLength + + ", ciphertextRead=" + ciphertextRead); + } + int bytesToRead = (int) Math.min((long) readableBytes, expectedRemaining); // The smallest ciphertext size is 16 bytes for the auth tag ciphertextBuffer.limit(ciphertextBuffer.position() + bytesToRead); ciphertextNettyBuf.readBytes(ciphertextBuffer); diff --git a/common/network-common/src/test/java/org/apache/spark/network/crypto/GcmAuthEngineSuite.java b/common/network-common/src/test/java/org/apache/spark/network/crypto/GcmAuthEngineSuite.java index baa2ac22398d7..43a7c67c20135 100644 --- a/common/network-common/src/test/java/org/apache/spark/network/crypto/GcmAuthEngineSuite.java +++ b/common/network-common/src/test/java/org/apache/spark/network/crypto/GcmAuthEngineSuite.java @@ -17,6 +17,9 @@ package org.apache.spark.network.crypto; +import com.google.common.primitives.Longs; +import com.google.crypto.tink.subtle.AesGcmHkdfStreaming; +import com.google.crypto.tink.subtle.StreamSegmentEncrypter; import io.netty.buffer.ByteBuf; import io.netty.buffer.Unpooled; import io.netty.channel.ChannelHandlerContext; @@ -570,6 +573,58 @@ public void testSplitLengthPrefix() throws Exception { } } + @Test + public void testCiphertextLengthLargerThanMaxInt() throws Exception { + TransportConf gcmConf = getConf(2, false); + try (AuthEngine client = new AuthEngine("appId", "secret", gcmConf); + AuthEngine server = new AuthEngine("appId", "secret", gcmConf)) { + AuthMessage clientChallenge = client.challenge(); + AuthMessage serverResponse = server.response(clientChallenge); + client.deriveSessionCipher(clientChallenge, serverResponse); + GcmTransportCipher cipher = (GcmTransportCipher) server.sessionCipher(); + GcmTransportCipher.DecryptionHandler decryptionHandler = cipher.getDecryptionHandler(); + AesGcmHkdfStreaming streaming = cipher.getAesGcmHkdfStreaming(); + + long expectedLength = (long) GcmTransportCipher.LENGTH_HEADER_BYTES + + streaming.getHeaderLength() + Integer.MAX_VALUE + 1L; + StreamSegmentEncrypter encrypter = streaming.newStreamSegmentEncrypter( + Longs.toByteArray(expectedLength)); + ByteBuffer header = encrypter.getHeader(); + ByteBuf ciphertext = Unpooled.buffer(GcmTransportCipher.LENGTH_HEADER_BYTES + + header.remaining() + 1) + .writeLong(expectedLength) + .writeBytes(header) + .writeByte(0); + ChannelHandlerContext ctx = mock(ChannelHandlerContext.class); + + decryptionHandler.channelRead(ctx, ciphertext); + + verify(ctx, never()).fireChannelRead(any()); + } + } + + @Test + public void testInvalidExpectedCiphertextLength() throws Exception { + TransportConf gcmConf = getConf(2, false); + try (AuthEngine client = new AuthEngine("appId", "secret", gcmConf); + AuthEngine server = new AuthEngine("appId", "secret", gcmConf)) { + AuthMessage clientChallenge = client.challenge(); + AuthMessage serverResponse = server.response(clientChallenge); + client.deriveSessionCipher(clientChallenge, serverResponse); + GcmTransportCipher cipher = (GcmTransportCipher) server.sessionCipher(); + GcmTransportCipher.DecryptionHandler decryptionHandler = cipher.getDecryptionHandler(); + long invalidLength = (long) GcmTransportCipher.LENGTH_HEADER_BYTES + + cipher.getAesGcmHkdfStreaming().getHeaderLength() - 1; + ByteBuf ciphertext = Unpooled.buffer(8).writeLong(invalidLength); + + IllegalStateException error = assertThrows( + IllegalStateException.class, + () -> decryptionHandler.channelRead(mock(ChannelHandlerContext.class), ciphertext)); + + assertEquals("Invalid expected ciphertext length: " + invalidLength, error.getMessage()); + } + } + /** * Regression test for the encryptedCount miscalculation that caused shuffle fetch stalls * for plaintext sizes in (plaintextSegmentSize - getCiphertextOffset(), plaintextSegmentSize]