From a58fdeaf3fc1b8edaa3f42a6679df133e318df0e 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 2864c9e6345..514e654211e 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 @@ -33,6 +33,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; @@ -202,8 +203,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 e52ad89f3ac..4de9d9d3e63 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 @@ -400,6 +400,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()); @@ -407,6 +410,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 8c741cbbf3e..63124e29b49 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 @@ -502,6 +502,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 5d912b471a0..1021d944f86 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.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.Arguments; @@ -44,9 +43,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.lang.Nullable; 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 cc8f396f6e3..7b6d494b2d8 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; @@ -48,12 +46,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();