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();