From 952198b2ee1f73ea8012a60a9968679e639be127 Mon Sep 17 00:00:00 2001 From: Brian Clozel Date: Tue, 28 Apr 2026 23:21:05 +0200 Subject: [PATCH] Fix PartGenerator token request while creating tmp file Prior to this commit, the `PartGenerator` would allow requesting additional part tokens while in the `CreateFileState`. This is invalid as any new token emitted would be rejected and would fail the entire process. This would only happen if the tmp file creation is slow enough for a new token to be parsed and emitted. This commit ensures that no new part token is requested while creating the temporary file. This change also fixes lifecycle issues and ensures that buffer resources are cleaned in case of errors. Fixes gh-36694 --- .../DefaultPartHttpMessageReader.java | 11 ++++-- .../http/codec/multipart/MultipartParser.java | 6 ++++ .../http/codec/multipart/PartGenerator.java | 5 +++ .../DefaultPartHttpMessageReaderTests.java | 34 +++++++++++++------ .../PartEventHttpMessageReaderTests.java | 8 ++--- 5 files changed, 46 insertions(+), 18 deletions(-) diff --git a/spring-web/src/main/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReader.java b/spring-web/src/main/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReader.java index 5038b5ba67f..5c38eb84c56 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReader.java +++ b/spring-web/src/main/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReader.java @@ -34,6 +34,7 @@ import reactor.core.scheduler.Schedulers; import org.springframework.core.ResolvableType; import org.springframework.core.codec.DecodingException; import org.springframework.core.io.buffer.DataBufferLimitException; +import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.http.MediaType; import org.springframework.http.ReactiveHttpInputMessage; import org.springframework.http.codec.HttpMessageReader; @@ -200,8 +201,14 @@ public class DefaultPartHttpMessageReader extends LoggingCodecSupport implements .windowUntil(MultipartParser.Token::isLast) .concatMap(partsTokens -> { if (tooManyParts(partCount)) { - return Mono.error(new DecodingException("Too many parts (" + partCount.get() + "/" + - this.maxParts + " allowed)")); + return partsTokens + .doOnNext(token -> { + if (token instanceof MultipartParser.BodyToken bodyToken) { + DataBufferUtils.release(bodyToken.buffer()); + } + }) + .then(Mono.error(new DecodingException("Too many parts (" + partCount.get() + "/" + + this.maxParts + " allowed)"))); } else { return PartGenerator.createPart(partsTokens, diff --git a/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartParser.java b/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartParser.java index 4448243609a..b3c8d1a10db 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartParser.java +++ b/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartParser.java @@ -401,6 +401,9 @@ final class MultipartParser extends BaseSubscriber { changeState(this, new BodyState(), buf); } + else { + changeState(this, DisposedState.INSTANCE, buf); + } } else { long count = this.byteCount.addAndGet(buf.readableByteCount()); @@ -408,6 +411,9 @@ final class MultipartParser extends BaseSubscriber { this.buffers.add(buf); requestBuffer(); } + else { + changeState(this, DisposedState.INSTANCE, buf); + } } } diff --git a/spring-web/src/main/java/org/springframework/http/codec/multipart/PartGenerator.java b/spring-web/src/main/java/org/springframework/http/codec/multipart/PartGenerator.java index 1c7ca5faee1..fa6074616b9 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/multipart/PartGenerator.java +++ b/spring-web/src/main/java/org/springframework/http/codec/multipart/PartGenerator.java @@ -500,6 +500,11 @@ final class PartGenerator extends BaseSubscriber { } } + @Override + public boolean canRequest() { + return false; + } + @Override public void dispose() { if (this.releaseOnDispose) { diff --git a/spring-web/src/test/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReaderTests.java b/spring-web/src/test/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReaderTests.java index c77ba81b33f..2212c33307d 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReaderTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/multipart/DefaultPartHttpMessageReaderTests.java @@ -28,7 +28,6 @@ import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.stream.Stream; -import io.netty.buffer.PooledByteBufAllocator; import org.jspecify.annotations.Nullable; import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; @@ -45,9 +44,9 @@ import org.springframework.core.codec.DecodingException; import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import org.springframework.core.io.buffer.DataBuffer; -import org.springframework.core.io.buffer.DataBufferFactory; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; -import org.springframework.core.io.buffer.NettyDataBufferFactory; +import org.springframework.core.testfixture.io.buffer.AbstractLeakCheckingTests; import org.springframework.http.MediaType; import org.springframework.web.testfixture.http.server.reactive.MockServerHttpRequest; @@ -60,9 +59,12 @@ import static org.springframework.core.ResolvableType.forClass; import static org.springframework.core.io.buffer.DataBufferUtils.release; /** + * Tests for {@link DefaultPartHttpMessageReader}. + * * @author Arjen Poutsma + * @author Brian Clozel */ -class DefaultPartHttpMessageReaderTests { +class DefaultPartHttpMessageReaderTests extends AbstractLeakCheckingTests { private static final String LOREM_IPSUM = "Lorem ipsum dolor sit amet, consectetur adipiscing elit. Integer iaculis metus id vestibulum nullam."; @@ -70,7 +72,6 @@ class DefaultPartHttpMessageReaderTests { private static final int BUFFER_SIZE = 64; - private static final DataBufferFactory bufferFactory = new NettyDataBufferFactory(new PooledByteBufAllocator()); @ParameterizedDefaultPartHttpMessageReaderTest void canRead(DefaultPartHttpMessageReader reader) { @@ -165,7 +166,7 @@ class DefaultPartHttpMessageReaderTests { new ClassPathResource("simple.multipart", getClass()), "simple-boundary"); Flux result = reader.read(forClass(Part.class), request, emptyMap()); - StepVerifier.create(result, 1) + StepVerifier.create(result) .consumeNextWith(part -> part.content().subscribe(DataBufferUtils::release)) .thenCancel() .verify(); @@ -221,7 +222,7 @@ class DefaultPartHttpMessageReaderTests { @Test void tooManyParts() throws InterruptedException { MockServerHttpRequest request = createRequest( - new ClassPathResource("simple.multipart", getClass()), "simple-boundary"); + new ClassPathResource("files.multipart", getClass()), "----WebKitFormBoundaryG8fJ50opQOML0oGD"); DefaultPartHttpMessageReader reader = new DefaultPartHttpMessageReader(); reader.setMaxParts(1); @@ -230,8 +231,7 @@ class DefaultPartHttpMessageReaderTests { CountDownLatch latch = new CountDownLatch(1); StepVerifier.create(result) - .consumeNextWith(part -> testPart(part, null, - "This is implicitly typed plain ASCII text.\r\nIt does NOT end with a linebreak.", latch)).as("Part 1") + .consumeNextWith(part -> testBrowserFile(part, "file2", "a.txt", LOREM_IPSUM, latch)).as("Part 1") .expectError(DecodingException.class) .verify(); @@ -276,7 +276,7 @@ class DefaultPartHttpMessageReaderTests { // gh-27612 @Test - void exceedHeaderLimit() throws InterruptedException { + void largeBufferForHeaderDoesNotExceedLimit() throws InterruptedException { Flux body = DataBufferUtils .readByteChannel((new ClassPathResource("files.multipart", getClass()))::readableChannel, bufferFactory, 282); @@ -300,6 +300,20 @@ class DefaultPartHttpMessageReaderTests { latch.await(); } + @Test + void exceedHeaderLimit() { + MockServerHttpRequest request = createRequest( + new ClassPathResource("files.multipart", getClass()), "\"----WebKitFormBoundaryG8fJ50opQOML0oGD\""); + + DefaultPartHttpMessageReader reader = new DefaultPartHttpMessageReader(); + reader.setMaxHeadersSize(80); + Flux result = reader.read(forClass(Part.class), request, emptyMap()); + + StepVerifier.create(result) + .expectError(DataBufferLimitException.class) + .verify(); + } + @ParameterizedDefaultPartHttpMessageReaderTest void emptyLastPart(DefaultPartHttpMessageReader reader) throws InterruptedException { MockServerHttpRequest request = createRequest( diff --git a/spring-web/src/test/java/org/springframework/http/codec/multipart/PartEventHttpMessageReaderTests.java b/spring-web/src/test/java/org/springframework/http/codec/multipart/PartEventHttpMessageReaderTests.java index 332b8c1d6be..70af1c9081a 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/multipart/PartEventHttpMessageReaderTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/multipart/PartEventHttpMessageReaderTests.java @@ -20,7 +20,6 @@ import java.nio.charset.StandardCharsets; import java.util.List; import java.util.function.Consumer; -import io.netty.buffer.PooledByteBufAllocator; import org.junit.jupiter.api.Test; import reactor.core.publisher.Flux; import reactor.test.StepVerifier; @@ -29,10 +28,9 @@ import org.springframework.core.codec.DecodingException; import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import org.springframework.core.io.buffer.DataBuffer; -import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; -import org.springframework.core.io.buffer.NettyDataBufferFactory; +import org.springframework.core.testfixture.io.buffer.AbstractLeakCheckingTests; import org.springframework.http.ContentDisposition; import org.springframework.http.HttpHeaders; import org.springframework.http.MediaType; @@ -47,12 +45,10 @@ import static org.springframework.core.ResolvableType.forClass; /** * @author Arjen Poutsma */ -class PartEventHttpMessageReaderTests { +class PartEventHttpMessageReaderTests extends AbstractLeakCheckingTests { private static final int BUFFER_SIZE = 64; - private static final DataBufferFactory bufferFactory = new NettyDataBufferFactory(new PooledByteBufAllocator()); - private static final MediaType TEXT_PLAIN_ASCII = new MediaType("text", "plain", StandardCharsets.US_ASCII); private final PartEventHttpMessageReader reader = new PartEventHttpMessageReader();