diff --git a/spring-core/src/main/java/org/springframework/core/codec/AbstractDataBufferDecoder.java b/spring-core/src/main/java/org/springframework/core/codec/AbstractDataBufferDecoder.java index b03d8079db..8fedf47576 100644 --- a/spring-core/src/main/java/org/springframework/core/codec/AbstractDataBufferDecoder.java +++ b/spring-core/src/main/java/org/springframework/core/codec/AbstractDataBufferDecoder.java @@ -47,12 +47,40 @@ import org.springframework.util.MimeType; */ public abstract class AbstractDataBufferDecoder extends AbstractDecoder { + private int maxInMemorySize = -1; + protected AbstractDataBufferDecoder(MimeType... supportedMimeTypes) { super(supportedMimeTypes); } + /** + * Configure a limit on the number of bytes that can be buffered whenever + * the input stream needs to be aggregated. This can be a result of + * decoding to a single {@code DataBuffer}, + * {@link java.nio.ByteBuffer ByteBuffer}, {@code byte[]}, + * {@link org.springframework.core.io.Resource Resource}, {@code String}, etc. + * It can also occur when splitting the input stream, e.g. delimited text, + * in which case the limit applies to data buffered between delimiters. + *

By default in 5.1 this is set to -1, unlimited. In 5.2 the default + * value for this limit is set to 256K. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + /** + * Return the {@link #setMaxInMemorySize configured} byte count limit. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + + @Override public Flux decode(Publisher input, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { @@ -64,7 +92,7 @@ public abstract class AbstractDataBufferDecoder extends AbstractDecoder { public Mono decodeToMono(Publisher input, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { - return DataBufferUtils.join(input) + return DataBufferUtils.join(input, this.maxInMemorySize) .map(buffer -> decodeDataBuffer(buffer, elementType, mimeType, hints)); } diff --git a/spring-core/src/main/java/org/springframework/core/codec/StringDecoder.java b/spring-core/src/main/java/org/springframework/core/codec/StringDecoder.java index 28cf7df55e..c3100490dd 100644 --- a/spring-core/src/main/java/org/springframework/core/codec/StringDecoder.java +++ b/spring-core/src/main/java/org/springframework/core/codec/StringDecoder.java @@ -25,14 +25,17 @@ import java.util.List; import java.util.Map; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; +import java.util.function.Consumer; import org.reactivestreams.Publisher; import reactor.core.publisher.Flux; import org.springframework.core.ResolvableType; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.core.io.buffer.DefaultDataBufferFactory; +import org.springframework.core.io.buffer.LimitedDataBufferList; import org.springframework.core.io.buffer.PooledDataBuffer; import org.springframework.core.log.LogFormatUtils; import org.springframework.lang.Nullable; @@ -92,10 +95,16 @@ public final class StringDecoder extends AbstractDataBufferDecoder { List delimiterBytes = getDelimiterBytes(mimeType); + // TODO: Drop Consumer and use bufferUntil with Supplier (reactor-core#1925) + // TODO: Drop doOnDiscard(LimitedDataBufferList.class, ...) (reactor-core#1924) + LimitedDataBufferConsumer limiter = new LimitedDataBufferConsumer(getMaxInMemorySize()); + Flux inputFlux = Flux.from(input) .flatMapIterable(buffer -> splitOnDelimiter(buffer, delimiterBytes)) + .doOnNext(limiter) .bufferUntil(buffer -> buffer == END_FRAME) .map(StringDecoder::joinUntilEndFrame) + .doOnDiscard(LimitedDataBufferList.class, LimitedDataBufferList::releaseAndClear) .doOnDiscard(PooledDataBuffer.class, DataBufferUtils::release); return super.decode(inputFlux, elementType, mimeType, hints); @@ -283,4 +292,35 @@ public final class StringDecoder extends AbstractDataBufferDecoder { new MimeType("text", "plain", DEFAULT_CHARSET), MimeTypeUtils.ALL); } + + /** + * Temporary measure for reactor-core#1925. + * Consumer that adds to a {@link LimitedDataBufferList} to enforce limits. + */ + private static class LimitedDataBufferConsumer implements Consumer { + + private final LimitedDataBufferList bufferList; + + + public LimitedDataBufferConsumer(int maxInMemorySize) { + this.bufferList = new LimitedDataBufferList(maxInMemorySize); + } + + + @Override + public void accept(DataBuffer buffer) { + if (buffer == END_FRAME) { + this.bufferList.clear(); + } + else { + try { + this.bufferList.add(buffer); + } + catch (DataBufferLimitException ex) { + DataBufferUtils.release(buffer); + throw ex; + } + } + } + } } diff --git a/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferLimitException.java b/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferLimitException.java new file mode 100644 index 0000000000..ee606aed57 --- /dev/null +++ b/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferLimitException.java @@ -0,0 +1,37 @@ +/* + * Copyright 2002-2019 the original author or authors. + * + * Licensed 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 + * + * https://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.springframework.core.io.buffer; + +/** + * Exception that indicates the cumulative number of bytes consumed from a + * stream of {@link DataBuffer DataBuffer}'s exceeded some pre-configured limit. + * This can be raised when data buffers are cached and aggregated, e.g. + * {@link DataBufferUtils#join}. Or it could also be raised when data buffers + * have been released but a parsed representation is being aggregated, e.g. async + * parsing with Jackson. + * + * @author Rossen Stoyanchev + * @since 5.1.11 + */ +@SuppressWarnings("serial") +public class DataBufferLimitException extends IllegalStateException { + + + public DataBufferLimitException(String message) { + super(message); + } + +} diff --git a/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferUtils.java b/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferUtils.java index 9ac2d80b77..c756efd735 100644 --- a/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferUtils.java +++ b/spring-core/src/main/java/org/springframework/core/io/buffer/DataBufferUtils.java @@ -437,14 +437,36 @@ public abstract class DataBufferUtils { * @since 5.0.3 */ public static Mono join(Publisher dataBuffers) { - Assert.notNull(dataBuffers, "'dataBuffers' must not be null"); + return join(dataBuffers, -1); + } - return Flux.from(dataBuffers) - .collectList() + /** + * Variant of {@link #join(Publisher)} that behaves the same way up until + * the specified max number of bytes to buffer. Once the limit is exceeded, + * {@link DataBufferLimitException} is raised. + * @param buffers the data buffers that are to be composed + * @param maxByteCount the max number of bytes to buffer, or -1 for unlimited + * @return a buffer with the aggregated content, possibly an empty Mono if + * the max number of bytes to buffer is exceeded. + * @throws DataBufferLimitException if maxByteCount is exceeded + * @since 5.1.11 + */ + @SuppressWarnings("unchecked") + public static Mono join(Publisher buffers, int maxByteCount) { + Assert.notNull(buffers, "'dataBuffers' must not be null"); + + if (buffers instanceof Mono) { + return (Mono) buffers; + } + + // TODO: Drop doOnDiscard(LimitedDataBufferList.class, ...) (reactor-core#1924) + + return Flux.from(buffers) + .collect(() -> new LimitedDataBufferList(maxByteCount), LimitedDataBufferList::add) .filter(list -> !list.isEmpty()) .map(list -> list.get(0).factory().join(list)) + .doOnDiscard(LimitedDataBufferList.class, LimitedDataBufferList::releaseAndClear) .doOnDiscard(PooledDataBuffer.class, DataBufferUtils::release); - } diff --git a/spring-core/src/main/java/org/springframework/core/io/buffer/LimitedDataBufferList.java b/spring-core/src/main/java/org/springframework/core/io/buffer/LimitedDataBufferList.java new file mode 100644 index 0000000000..fb8c42aeeb --- /dev/null +++ b/spring-core/src/main/java/org/springframework/core/io/buffer/LimitedDataBufferList.java @@ -0,0 +1,157 @@ +/* + * Copyright 2002-2019 the original author or authors. + * + * Licensed 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 + * + * https://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.springframework.core.io.buffer; + +import java.util.ArrayList; +import java.util.Collection; +import java.util.List; +import java.util.function.Predicate; + +import reactor.core.publisher.Flux; + +/** + * Custom {@link List} to collect data buffers with and enforce a + * limit on the total number of bytes buffered. For use with "collect" or + * other buffering operators in declarative APIs, e.g. {@link Flux}. + * + *

Adding elements increases the byte count and if the limit is exceeded, + * {@link DataBufferLimitException} is raised. {@link #clear()} resets the + * count. Remove and set are not supported. + * + *

Note: This class does not automatically release the + * buffers it contains. It is usually preferable to use hooks such as + * {@link Flux#doOnDiscard} that also take care of cancel and error signals, + * or otherwise {@link #releaseAndClear()} can be used. + * + * @author Rossen Stoyanchev + * @since 5.1.11 + */ +@SuppressWarnings("serial") +public class LimitedDataBufferList extends ArrayList { + + private final int maxByteCount; + + private int byteCount; + + + public LimitedDataBufferList(int maxByteCount) { + this.maxByteCount = maxByteCount; + } + + + @Override + public boolean add(DataBuffer buffer) { + boolean result = super.add(buffer); + if (result) { + updateCount(buffer.readableByteCount()); + } + return result; + } + + @Override + public void add(int index, DataBuffer buffer) { + super.add(index, buffer); + updateCount(buffer.readableByteCount()); + } + + @Override + public boolean addAll(Collection collection) { + boolean result = super.addAll(collection); + collection.forEach(buffer -> updateCount(buffer.readableByteCount())); + return result; + } + + @Override + public boolean addAll(int index, Collection collection) { + boolean result = super.addAll(index, collection); + collection.forEach(buffer -> updateCount(buffer.readableByteCount())); + return result; + } + + private void updateCount(int bytesToAdd) { + if (this.maxByteCount < 0) { + return; + } + if (bytesToAdd > Integer.MAX_VALUE - this.byteCount) { + raiseLimitException(); + } + else { + this.byteCount += bytesToAdd; + if (this.byteCount > this.maxByteCount) { + raiseLimitException(); + } + } + } + + private void raiseLimitException() { + // Do not release here, it's likely down via doOnDiscard.. + throw new DataBufferLimitException( + "Exceeded limit on max bytes to buffer : " + this.maxByteCount); + } + + @Override + public DataBuffer remove(int index) { + throw new UnsupportedOperationException(); + } + + @Override + public boolean remove(Object o) { + throw new UnsupportedOperationException(); + } + + @Override + protected void removeRange(int fromIndex, int toIndex) { + throw new UnsupportedOperationException(); + } + + @Override + public boolean removeAll(Collection c) { + throw new UnsupportedOperationException(); + } + + @Override + public boolean removeIf(Predicate filter) { + throw new UnsupportedOperationException(); + } + + @Override + public DataBuffer set(int index, DataBuffer element) { + throw new UnsupportedOperationException(); + } + + @Override + public void clear() { + this.byteCount = 0; + super.clear(); + } + + /** + * Shortcut to {@link DataBufferUtils#release release} all data buffers and + * then {@link #clear()}. + */ + public void releaseAndClear() { + forEach(buf -> { + try { + DataBufferUtils.release(buf); + } + catch (Throwable ex) { + // Keep going.. + } + }); + clear(); + } + +} diff --git a/spring-core/src/test/java/org/springframework/core/codec/StringDecoderTests.java b/spring-core/src/test/java/org/springframework/core/codec/StringDecoderTests.java index ed10c96c83..b584080af2 100644 --- a/spring-core/src/test/java/org/springframework/core/codec/StringDecoderTests.java +++ b/spring-core/src/test/java/org/springframework/core/codec/StringDecoderTests.java @@ -29,6 +29,7 @@ import reactor.test.StepVerifier; import org.springframework.core.ResolvableType; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.util.MimeType; import org.springframework.util.MimeTypeUtils; @@ -126,6 +127,20 @@ public class StringDecoderTests extends AbstractDecoderTestCase { .verify()); } + @Test + public void decodeNewLineWithLimit() { + Flux input = Flux.just( + stringBuffer("abc\n"), + stringBuffer("defg\n"), + stringBuffer("hijkl\n") + ); + this.decoder.setMaxInMemorySize(4); + + testDecode(input, String.class, step -> + step.expectNext("abc", "defg") + .verifyError(DataBufferLimitException.class)); + } + @Test public void decodeNewLineIncludeDelimiters() { this.decoder = StringDecoder.allMimeTypes(StringDecoder.DEFAULT_DELIMITERS, false); diff --git a/spring-core/src/test/java/org/springframework/core/io/buffer/DataBufferUtilsTests.java b/spring-core/src/test/java/org/springframework/core/io/buffer/DataBufferUtilsTests.java index 115d0fce6b..456c588166 100644 --- a/spring-core/src/test/java/org/springframework/core/io/buffer/DataBufferUtilsTests.java +++ b/spring-core/src/test/java/org/springframework/core/io/buffer/DataBufferUtilsTests.java @@ -48,11 +48,14 @@ import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.Resource; import org.springframework.core.io.buffer.support.DataBufferTestUtils; -import static org.junit.Assert.*; -import static org.mockito.ArgumentMatchers.*; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.fail; +import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.anyLong; +import static org.mockito.Mockito.doAnswer; import static org.mockito.Mockito.isA; -import static org.mockito.Mockito.*; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; /** * @author Arjen Poutsma @@ -716,14 +719,25 @@ public class DataBufferUtilsTests extends AbstractDataBufferAllocatingTestCase { Mono result = DataBufferUtils.join(flux); StepVerifier.create(result) - .consumeNextWith(dataBuffer -> { - assertEquals("foobarbaz", - DataBufferTestUtils.dumpString(dataBuffer, StandardCharsets.UTF_8)); - release(dataBuffer); + .consumeNextWith(buf -> { + assertEquals("foobarbaz", DataBufferTestUtils.dumpString(buf, StandardCharsets.UTF_8)); + release(buf); }) .verifyComplete(); } + @Test + public void joinWithLimit() { + DataBuffer foo = stringBuffer("foo"); + DataBuffer bar = stringBuffer("bar"); + DataBuffer baz = stringBuffer("baz"); + Flux flux = Flux.just(foo, bar, baz); + Mono result = DataBufferUtils.join(flux, 8); + + StepVerifier.create(result) + .verifyError(DataBufferLimitException.class); + } + @Test public void joinErrors() { DataBuffer foo = stringBuffer("foo"); diff --git a/spring-core/src/test/java/org/springframework/core/io/buffer/LimitedDataBufferListTests.java b/spring-core/src/test/java/org/springframework/core/io/buffer/LimitedDataBufferListTests.java new file mode 100644 index 0000000000..baf348e5f9 --- /dev/null +++ b/spring-core/src/test/java/org/springframework/core/io/buffer/LimitedDataBufferListTests.java @@ -0,0 +1,63 @@ +/* + * Copyright 2002-2019 the original author or authors. + * + * Licensed 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 + * + * https://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.springframework.core.io.buffer; + +import java.nio.charset.StandardCharsets; + +import org.junit.Test; + +import static org.junit.Assert.fail; + +/** + * Unit tests for {@link LimitedDataBufferList}. + * @author Rossen Stoyanchev + * @since 5.1.11 + */ +public class LimitedDataBufferListTests { + + private final static DataBufferFactory bufferFactory = new DefaultDataBufferFactory(); + + + @Test + public void limitEnforced() { + try { + new LimitedDataBufferList(5).add(toDataBuffer("123456")); + fail(); + } + catch (DataBufferLimitException ex) { + // Expected + } + } + + @Test + public void limitIgnored() { + new LimitedDataBufferList(-1).add(toDataBuffer("123456")); + } + + @Test + public void clearResetsCount() { + LimitedDataBufferList list = new LimitedDataBufferList(5); + list.add(toDataBuffer("12345")); + list.clear(); + list.add(toDataBuffer("12345")); + } + + + private static DataBuffer toDataBuffer(String value) { + return bufferFactory.wrap(value.getBytes(StandardCharsets.UTF_8)); + } + +} diff --git a/spring-web/src/main/java/org/springframework/http/codec/CodecConfigurer.java b/spring-web/src/main/java/org/springframework/http/codec/CodecConfigurer.java index 17b820a930..40a18a08d8 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/CodecConfigurer.java +++ b/spring-web/src/main/java/org/springframework/http/codec/CodecConfigurer.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -143,6 +143,21 @@ public interface CodecConfigurer { */ void jaxb2Encoder(Encoder encoder); + /** + * Configure a limit on the number of bytes that can be buffered whenever + * the input stream needs to be aggregated. This can be a result of + * decoding to a single {@code DataBuffer}, + * {@link java.nio.ByteBuffer ByteBuffer}, {@code byte[]}, + * {@link org.springframework.core.io.Resource Resource}, {@code String}, etc. + * It can also occur when splitting the input stream, e.g. delimited text, + * in which case the limit applies to data buffered between delimiters. + *

By default this is not set, in which case individual codec defaults + * apply. In 5.1 most codecs are not limited except {@code FormHttpMessageReader} + * which is limited to 256K. In 5.2 all codecs are limited to 256K by default. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @sine 5.1.11 + */ + void maxInMemorySize(int byteCount); /** * Whether to log form data at DEBUG level, and headers at TRACE level. * Both may contain sensitive information. diff --git a/spring-web/src/main/java/org/springframework/http/codec/FormHttpMessageReader.java b/spring-web/src/main/java/org/springframework/http/codec/FormHttpMessageReader.java index 01c2d30b33..39c75a8578 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/FormHttpMessageReader.java +++ b/spring-web/src/main/java/org/springframework/http/codec/FormHttpMessageReader.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -30,6 +30,7 @@ import reactor.core.publisher.Mono; import org.springframework.core.ResolvableType; import org.springframework.core.codec.Hints; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.core.log.LogFormatUtils; import org.springframework.http.MediaType; @@ -62,6 +63,8 @@ public class FormHttpMessageReader extends LoggingCodecSupport private Charset defaultCharset = DEFAULT_CHARSET; + private int maxInMemorySize = 256 * 1024; + /** * Set the default character set to use for reading form data when the @@ -80,6 +83,26 @@ public class FormHttpMessageReader extends LoggingCodecSupport return this.defaultCharset; } + /** + * Set the max number of bytes for input form data. As form data is buffered + * before it is parsed, this helps to limit the amount of buffering. Once + * the limit is exceeded, {@link DataBufferLimitException} is raised. + *

By default this is set to 256K. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + /** + * Return the {@link #setMaxInMemorySize configured} byte count limit. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + @Override public boolean canRead(ResolvableType elementType, @Nullable MediaType mediaType) { @@ -105,7 +128,7 @@ public class FormHttpMessageReader extends LoggingCodecSupport MediaType contentType = message.getHeaders().getContentType(); Charset charset = getMediaTypeCharset(contentType); - return DataBufferUtils.join(message.getBody()) + return DataBufferUtils.join(message.getBody(), getMaxInMemorySize()) .map(buffer -> { CharBuffer charBuffer = charset.decode(buffer.asByteBuffer()); String body = charBuffer.toString(); diff --git a/spring-web/src/main/java/org/springframework/http/codec/ServerCodecConfigurer.java b/spring-web/src/main/java/org/springframework/http/codec/ServerCodecConfigurer.java index 9e4350e72e..59a209ac59 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/ServerCodecConfigurer.java +++ b/spring-web/src/main/java/org/springframework/http/codec/ServerCodecConfigurer.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -76,6 +76,20 @@ public interface ServerCodecConfigurer extends CodecConfigurer { */ interface ServerDefaultCodecs extends DefaultCodecs { + /** + * Configure the {@code HttpMessageReader} to use for multipart requests. + *

By default, if + * Synchronoss NIO Multipart + * is present, this is set to + * {@link org.springframework.http.codec.multipart.MultipartHttpMessageReader + * MultipartHttpMessageReader} created with an instance of + * {@link org.springframework.http.codec.multipart.SynchronossPartHttpMessageReader + * SynchronossPartHttpMessageReader}. + * @param reader the message reader to use for multipart requests. + * @since 5.1.11 + */ + void multipartReader(HttpMessageReader reader); + /** * Configure the {@code Encoder} to use for Server-Sent Events. *

By default if this is not set, and Jackson is available, the diff --git a/spring-web/src/main/java/org/springframework/http/codec/json/AbstractJackson2Decoder.java b/spring-web/src/main/java/org/springframework/http/codec/json/AbstractJackson2Decoder.java index b42e988d54..95586080ae 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/json/AbstractJackson2Decoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/json/AbstractJackson2Decoder.java @@ -38,6 +38,7 @@ import org.springframework.core.codec.CodecException; import org.springframework.core.codec.DecodingException; import org.springframework.core.codec.Hints; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.log.LogFormatUtils; import org.springframework.http.codec.HttpMessageDecoder; import org.springframework.http.server.reactive.ServerHttpRequest; @@ -57,6 +58,9 @@ import org.springframework.util.MimeType; */ public abstract class AbstractJackson2Decoder extends Jackson2CodecSupport implements HttpMessageDecoder { + private int maxInMemorySize = -1; + + /** * Until https://github.com/FasterXML/jackson-core/issues/476 is resolved, * we need to ensure buffer recycling is off. @@ -74,6 +78,29 @@ public abstract class AbstractJackson2Decoder extends Jackson2CodecSupport imple } + /** + * Set the max number of bytes that can be buffered by this decoder. This + * is either the size of the entire input when decoding as a whole, or the + * size of one top-level JSON object within a JSON stream. When the limit + * is exceeded, {@link DataBufferLimitException} is raised. + *

By default in 5.1 this is set to -1, unlimited. In 5.2 the default + * value for this limit is set to 256K. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + /** + * Return the {@link #setMaxInMemorySize configured} byte count limit. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + + @Override public boolean canDecode(ResolvableType elementType, @Nullable MimeType mimeType) { JavaType javaType = getObjectMapper().getTypeFactory().constructType(elementType.getType()); @@ -87,7 +114,7 @@ public abstract class AbstractJackson2Decoder extends Jackson2CodecSupport imple @Nullable MimeType mimeType, @Nullable Map hints) { Flux tokens = Jackson2Tokenizer.tokenize( - Flux.from(input), this.jsonFactory, getObjectMapper(), true); + Flux.from(input), this.jsonFactory, getObjectMapper(), true, getMaxInMemorySize()); return decodeInternal(tokens, elementType, mimeType, hints); } @@ -96,7 +123,7 @@ public abstract class AbstractJackson2Decoder extends Jackson2CodecSupport imple @Nullable MimeType mimeType, @Nullable Map hints) { Flux tokens = Jackson2Tokenizer.tokenize( - Flux.from(input), this.jsonFactory, getObjectMapper(), false); + Flux.from(input), this.jsonFactory, getObjectMapper(), false, getMaxInMemorySize()); return decodeInternal(tokens, elementType, mimeType, hints).singleOrEmpty(); } diff --git a/spring-web/src/main/java/org/springframework/http/codec/json/Jackson2Tokenizer.java b/spring-web/src/main/java/org/springframework/http/codec/json/Jackson2Tokenizer.java index 8549dee0bd..606f0a34aa 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/json/Jackson2Tokenizer.java +++ b/spring-web/src/main/java/org/springframework/http/codec/json/Jackson2Tokenizer.java @@ -34,6 +34,7 @@ import reactor.core.publisher.Flux; import org.springframework.core.codec.DecodingException; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; /** @@ -60,30 +61,39 @@ final class Jackson2Tokenizer { private int arrayDepth; + private final int maxInMemorySize; + + private int byteCount; + + // TODO: change to ByteBufferFeeder when supported by Jackson // See https://github.com/FasterXML/jackson-core/issues/478 private final ByteArrayFeeder inputFeeder; - private Jackson2Tokenizer( - JsonParser parser, DeserializationContext deserializationContext, boolean tokenizeArrayElements) { + private Jackson2Tokenizer(JsonParser parser, DeserializationContext deserializationContext, + boolean tokenizeArrayElements, int maxInMemorySize) { this.parser = parser; this.deserializationContext = deserializationContext; this.tokenizeArrayElements = tokenizeArrayElements; this.tokenBuffer = new TokenBuffer(parser, deserializationContext); this.inputFeeder = (ByteArrayFeeder) this.parser.getNonBlockingInputFeeder(); + this.maxInMemorySize = maxInMemorySize; } private Flux tokenize(DataBuffer dataBuffer) { + int bufferSize = dataBuffer.readableByteCount(); byte[] bytes = new byte[dataBuffer.readableByteCount()]; dataBuffer.read(bytes); DataBufferUtils.release(dataBuffer); try { this.inputFeeder.feedInput(bytes, 0, bytes.length); - return parseTokenBufferFlux(); + List result = parseTokenBufferFlux(); + assertInMemorySize(bufferSize, result); + return Flux.fromIterable(result); } catch (JsonProcessingException ex) { return Flux.error(new DecodingException("JSON decoding error: " + ex.getOriginalMessage(), ex)); @@ -96,7 +106,8 @@ final class Jackson2Tokenizer { private Flux endOfInput() { this.inputFeeder.endOfInput(); try { - return parseTokenBufferFlux(); + List result = parseTokenBufferFlux(); + return Flux.fromIterable(result); } catch (JsonProcessingException ex) { return Flux.error(new DecodingException("JSON decoding error: " + ex.getOriginalMessage(), ex)); @@ -106,7 +117,7 @@ final class Jackson2Tokenizer { } } - private Flux parseTokenBufferFlux() throws IOException { + private List parseTokenBufferFlux() throws IOException { List result = new ArrayList<>(); while (true) { @@ -124,7 +135,7 @@ final class Jackson2Tokenizer { processTokenArray(token, result); } } - return Flux.fromIterable(result); + return result; } private void updateDepth(JsonToken token) { @@ -171,18 +182,40 @@ final class Jackson2Tokenizer { (token == JsonToken.END_ARRAY && this.arrayDepth == 0)); } + private void assertInMemorySize(int currentBufferSize, List result) { + if (this.maxInMemorySize >= 0) { + if (!result.isEmpty()) { + this.byteCount = 0; + } + else if (currentBufferSize > Integer.MAX_VALUE - this.byteCount) { + raiseLimitException(); + } + else { + this.byteCount += currentBufferSize; + if (this.byteCount > this.maxInMemorySize) { + raiseLimitException(); + } + } + } + } + + private void raiseLimitException() { + throw new DataBufferLimitException( + "Exceeded limit on max bytes per JSON object: " + this.maxInMemorySize); + } + /** * Tokenize the given {@code Flux} into {@code Flux}. * @param dataBuffers the source data buffers * @param jsonFactory the factory to use * @param objectMapper the current mapper instance - * @param tokenizeArrayElements if {@code true} and the "top level" JSON object is + * @param tokenizeArrays if {@code true} and the "top level" JSON object is * an array, each element is returned individually immediately after it is received * @return the resulting token buffers */ public static Flux tokenize(Flux dataBuffers, JsonFactory jsonFactory, - ObjectMapper objectMapper, boolean tokenizeArrayElements) { + ObjectMapper objectMapper, boolean tokenizeArrays, int maxInMemorySize) { try { JsonParser parser = jsonFactory.createNonBlockingByteArrayParser(); @@ -191,7 +224,7 @@ final class Jackson2Tokenizer { context = ((DefaultDeserializationContext) context).createInstance( objectMapper.getDeserializationConfig(), parser, objectMapper.getInjectableValues()); } - Jackson2Tokenizer tokenizer = new Jackson2Tokenizer(parser, context, tokenizeArrayElements); + Jackson2Tokenizer tokenizer = new Jackson2Tokenizer(parser, context, tokenizeArrays, maxInMemorySize); return dataBuffers.flatMap(tokenizer::tokenize, Flux::error, tokenizer::endOfInput); } catch (IOException ex) { diff --git a/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartHttpMessageReader.java b/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartHttpMessageReader.java index ffecd7c518..3c8c4b483e 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartHttpMessageReader.java +++ b/spring-web/src/main/java/org/springframework/http/codec/multipart/MultipartHttpMessageReader.java @@ -65,6 +65,14 @@ public class MultipartHttpMessageReader extends LoggingCodecSupport } + /** + * Return the configured parts reader. + * @since 5.1.11 + */ + public HttpMessageReader getPartReader() { + return this.partReader; + } + @Override public List getReadableMediaTypes() { return Collections.singletonList(MediaType.MULTIPART_FORM_DATA); diff --git a/spring-web/src/main/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReader.java b/spring-web/src/main/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReader.java index 642e16688f..855b9fb047 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReader.java +++ b/spring-web/src/main/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReader.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -40,14 +40,18 @@ import org.synchronoss.cloud.nio.multipart.NioMultipartParser; import org.synchronoss.cloud.nio.multipart.NioMultipartParserListener; import org.synchronoss.cloud.nio.multipart.PartBodyStreamStorageFactory; import org.synchronoss.cloud.nio.stream.storage.StreamStorage; +import reactor.core.publisher.BaseSubscriber; import reactor.core.publisher.Flux; import reactor.core.publisher.FluxSink; import reactor.core.publisher.Mono; +import reactor.core.publisher.SignalType; import org.springframework.core.ResolvableType; +import org.springframework.core.codec.DecodingException; import org.springframework.core.codec.Hints; 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.DefaultDataBufferFactory; import org.springframework.core.log.LogFormatUtils; @@ -69,15 +73,83 @@ import org.springframework.util.Assert; * @author Sebastien Deleuze * @author Rossen Stoyanchev * @author Arjen Poutsma + * @author Brian Clozel * @since 5.0 * @see Synchronoss NIO Multipart * @see MultipartHttpMessageReader */ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implements HttpMessageReader { - private final DataBufferFactory bufferFactory = new DefaultDataBufferFactory(); + // Static DataBufferFactory to copy from FileInputStream or wrap bytes[]. + private static final DataBufferFactory bufferFactory = new DefaultDataBufferFactory(); - private final PartBodyStreamStorageFactory streamStorageFactory = new DefaultPartBodyStreamStorageFactory(); + + private int maxInMemorySize = -1; + + private long maxDiskUsagePerPart = -1; + + private long maxParts = -1; + + + /** + * Configure the maximum amount of memory that is allowed to use per part. + * When the limit is exceeded: + *

    + *
  • file parts are written to a temporary file. + *
  • non-file parts are rejected with {@link DataBufferLimitException}. + *
+ *

By default in 5.1 this is set to -1 in which case this limit is + * not enforced and all parts may be written to disk and are limited only + * by the {@link #setMaxDiskUsagePerPart(long) maxDiskUsagePerPart} property. + * In 5.2 this default value for this limit is set to 256K. + * @param byteCount the in-memory limit in bytes, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + /** + * Get the {@link #setMaxInMemorySize configured} maximum in-memory size. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + + /** + * Configure the maximum amount of disk space allowed for file parts. + *

By default this is set to -1. + * @param maxDiskUsagePerPart the disk limit in bytes, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxDiskUsagePerPart(long maxDiskUsagePerPart) { + this.maxDiskUsagePerPart = maxDiskUsagePerPart; + } + + /** + * Get the {@link #setMaxDiskUsagePerPart configured} maximum disk usage. + * @since 5.1.11 + */ + public long getMaxDiskUsagePerPart() { + return this.maxDiskUsagePerPart; + } + + /** + * Specify the maximum number of parts allowed in a given multipart request. + * @since 5.1.11 + */ + public void setMaxParts(long maxParts) { + this.maxParts = maxParts; + } + + /** + * Return the {@link #setMaxParts configured} limit on the number of parts. + * @since 5.1.11 + */ + public long getMaxParts() { + return this.maxParts; + } @Override @@ -94,7 +166,7 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem @Override public Flux read(ResolvableType elementType, ReactiveHttpInputMessage message, Map hints) { - return Flux.create(new SynchronossPartGenerator(message, this.bufferFactory, this.streamStorageFactory)) + return Flux.create(new SynchronossPartGenerator(message)) .doOnNext(part -> { if (!Hints.isLoggingSuppressed(hints)) { LogFormatUtils.traceDebug(logger, traceOn -> Hints.getLogPrefix(hints) + "Parsed " + @@ -107,33 +179,36 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem @Override - public Mono readMono(ResolvableType elementType, ReactiveHttpInputMessage message, Map hints) { - return Mono.error(new UnsupportedOperationException("Cannot read multipart request body into single Part")); + public Mono readMono( + ResolvableType elementType, ReactiveHttpInputMessage message, Map hints) { + + return Mono.error(new UnsupportedOperationException( + "Cannot read multipart request body into single Part")); } /** - * Consume and feed input to the Synchronoss parser, then listen for parser - * output events and adapt to {@code Flux>}. + * Subscribe to the input stream and feed the Synchronoss parser. Then listen + * for parser output, creating parts, and pushing them into the FluxSink. */ - private static class SynchronossPartGenerator implements Consumer> { + private class SynchronossPartGenerator extends BaseSubscriber implements Consumer> { private final ReactiveHttpInputMessage inputMessage; - private final DataBufferFactory bufferFactory; + private final LimitedPartBodyStreamStorageFactory storageFactory = new LimitedPartBodyStreamStorageFactory(); - private final PartBodyStreamStorageFactory streamStorageFactory; + private NioMultipartParserListener listener; - SynchronossPartGenerator(ReactiveHttpInputMessage inputMessage, DataBufferFactory bufferFactory, - PartBodyStreamStorageFactory streamStorageFactory) { + private NioMultipartParser parser; + + public SynchronossPartGenerator(ReactiveHttpInputMessage inputMessage) { this.inputMessage = inputMessage; - this.bufferFactory = bufferFactory; - this.streamStorageFactory = streamStorageFactory; } + @Override - public void accept(FluxSink emitter) { + public void accept(FluxSink sink) { HttpHeaders headers = this.inputMessage.getHeaders(); MediaType mediaType = headers.getContentType(); Assert.state(mediaType != null, "No content type set"); @@ -142,40 +217,57 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem Charset charset = Optional.ofNullable(mediaType.getCharset()).orElse(StandardCharsets.UTF_8); MultipartContext context = new MultipartContext(mediaType.toString(), length, charset.name()); - NioMultipartParserListener listener = new FluxSinkAdapterListener(emitter, this.bufferFactory, context); - NioMultipartParser parser = Multipart - .multipart(context) - .usePartBodyStreamStorageFactory(this.streamStorageFactory) - .forNIO(listener); + this.listener = new FluxSinkAdapterListener(sink, context, this.storageFactory); - this.inputMessage.getBody().subscribe(buffer -> { - byte[] resultBytes = new byte[buffer.readableByteCount()]; - buffer.read(resultBytes); - try { - parser.write(resultBytes); - } - catch (IOException ex) { - listener.onError("Exception thrown providing input to the parser", ex); - } - finally { - DataBufferUtils.release(buffer); - } - }, ex -> { - try { - listener.onError("Request body input error", ex); - parser.close(); - } - catch (IOException ex2) { - listener.onError("Exception thrown while closing the parser", ex2); - } - }, () -> { - try { - parser.close(); - } - catch (IOException ex) { - listener.onError("Exception thrown while closing the parser", ex); - } - }); + this.parser = Multipart + .multipart(context) + .usePartBodyStreamStorageFactory(this.storageFactory) + .forNIO(this.listener); + + this.inputMessage.getBody().subscribe(this); + } + + @Override + protected void hookOnNext(DataBuffer buffer) { + int size = buffer.readableByteCount(); + this.storageFactory.increaseByteCount(size); + byte[] resultBytes = new byte[size]; + buffer.read(resultBytes); + try { + this.parser.write(resultBytes); + } + catch (IOException ex) { + cancel(); + int index = this.storageFactory.getCurrentPartIndex(); + this.listener.onError("Parser error for part [" + index + "]", ex); + } + finally { + DataBufferUtils.release(buffer); + } + } + + @Override + protected void hookOnError(Throwable ex) { + try { + this.parser.close(); + } + catch (IOException ex2) { + // ignore + } + finally { + int index = this.storageFactory.getCurrentPartIndex(); + this.listener.onError("Failure while parsing part[" + index + "]", ex); + } + } + + @Override + protected void hookFinally(SignalType type) { + try { + this.parser.close(); + } + catch (IOException ex) { + this.listener.onError("Error while closing parser", ex); + } } private int getContentLength(HttpHeaders headers) { @@ -186,6 +278,54 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem } + private class LimitedPartBodyStreamStorageFactory implements PartBodyStreamStorageFactory { + + private final PartBodyStreamStorageFactory storageFactory = maxInMemorySize > 0 ? + new DefaultPartBodyStreamStorageFactory(maxInMemorySize) : + new DefaultPartBodyStreamStorageFactory(); + + private int index = 1; + + private boolean isFilePart; + + private long partSize; + + + public int getCurrentPartIndex() { + return this.index; + } + + @Override + public StreamStorage newStreamStorageForPartBody(Map> headers, int index) { + this.index = index; + this.isFilePart = (MultipartUtils.getFileName(headers) != null); + this.partSize = 0; + if (maxParts > 0 && index > maxParts) { + throw new DecodingException("Too many parts (" + index + " allowed)"); + } + return this.storageFactory.newStreamStorageForPartBody(headers, index); + } + + public void increaseByteCount(long byteCount) { + this.partSize += byteCount; + if (maxInMemorySize > 0 && !this.isFilePart && this.partSize >= maxInMemorySize) { + throw new DataBufferLimitException("Part[" + this.index + "] " + + "exceeded the in-memory limit of " + maxInMemorySize + " bytes"); + } + if (maxDiskUsagePerPart > 0 && this.isFilePart && this.partSize > maxDiskUsagePerPart) { + throw new DecodingException("Part[" + this.index + "] " + + "exceeded the disk usage limit of " + maxDiskUsagePerPart + " bytes"); + } + } + + public void partFinished() { + this.index++; + this.isFilePart = false; + this.partSize = 0; + } + } + + /** * Listen for parser output and adapt to {@code Flux>}. */ @@ -193,43 +333,48 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem private final FluxSink sink; - private final DataBufferFactory bufferFactory; - private final MultipartContext context; + private final LimitedPartBodyStreamStorageFactory storageFactory; + private final AtomicInteger terminated = new AtomicInteger(0); - FluxSinkAdapterListener(FluxSink sink, DataBufferFactory factory, MultipartContext context) { + + FluxSinkAdapterListener( + FluxSink sink, MultipartContext context, LimitedPartBodyStreamStorageFactory factory) { + this.sink = sink; - this.bufferFactory = factory; this.context = context; + this.storageFactory = factory; } + @Override public void onPartFinished(StreamStorage storage, Map> headers) { HttpHeaders httpHeaders = new HttpHeaders(); httpHeaders.putAll(headers); + this.storageFactory.partFinished(); this.sink.next(createPart(storage, httpHeaders)); } private Part createPart(StreamStorage storage, HttpHeaders httpHeaders) { String filename = MultipartUtils.getFileName(httpHeaders); if (filename != null) { - return new SynchronossFilePart(httpHeaders, filename, storage, this.bufferFactory); + return new SynchronossFilePart(httpHeaders, filename, storage); } else if (MultipartUtils.isFormField(httpHeaders, this.context)) { String value = MultipartUtils.readFormParameterValue(storage, httpHeaders); - return new SynchronossFormFieldPart(httpHeaders, this.bufferFactory, value); + return new SynchronossFormFieldPart(httpHeaders, value); } else { - return new SynchronossPart(httpHeaders, storage, this.bufferFactory); + return new SynchronossPart(httpHeaders, storage); } } @Override public void onError(String message, Throwable cause) { if (this.terminated.getAndIncrement() == 0) { - this.sink.error(new RuntimeException(message, cause)); + this.sink.error(new DecodingException(message, cause)); } } @@ -256,14 +401,10 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem private final HttpHeaders headers; - private final DataBufferFactory bufferFactory; - - AbstractSynchronossPart(HttpHeaders headers, DataBufferFactory bufferFactory) { + AbstractSynchronossPart(HttpHeaders headers) { Assert.notNull(headers, "HttpHeaders is required"); - Assert.notNull(bufferFactory, "DataBufferFactory is required"); this.name = MultipartUtils.getFieldName(headers); this.headers = headers; - this.bufferFactory = bufferFactory; } @Override @@ -276,10 +417,6 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem return this.headers; } - DataBufferFactory getBufferFactory() { - return this.bufferFactory; - } - @Override public String toString() { return "Part '" + this.name + "', headers=" + this.headers; @@ -291,15 +428,15 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem private final StreamStorage storage; - SynchronossPart(HttpHeaders headers, StreamStorage storage, DataBufferFactory factory) { - super(headers, factory); + SynchronossPart(HttpHeaders headers, StreamStorage storage) { + super(headers); Assert.notNull(storage, "StreamStorage is required"); this.storage = storage; } @Override public Flux content() { - return DataBufferUtils.readInputStream(getStorage()::getInputStream, getBufferFactory(), 4096); + return DataBufferUtils.readInputStream(getStorage()::getInputStream, bufferFactory, 4096); } protected StreamStorage getStorage() { @@ -315,8 +452,8 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem private final String filename; - SynchronossFilePart(HttpHeaders headers, String filename, StreamStorage storage, DataBufferFactory factory) { - super(headers, storage, factory); + SynchronossFilePart(HttpHeaders headers, String filename, StreamStorage storage) { + super(headers, storage); this.filename = filename; } @@ -375,8 +512,8 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem private final String content; - SynchronossFormFieldPart(HttpHeaders headers, DataBufferFactory bufferFactory, String content) { - super(headers, bufferFactory); + SynchronossFormFieldPart(HttpHeaders headers, String content) { + super(headers); this.content = content; } @@ -388,9 +525,7 @@ public class SynchronossPartHttpMessageReader extends LoggingCodecSupport implem @Override public Flux content() { byte[] bytes = this.content.getBytes(getCharset()); - DataBuffer buffer = getBufferFactory().allocateBuffer(bytes.length); - buffer.write(bytes); - return Flux.just(buffer); + return Flux.just(bufferFactory.wrap(bytes)); } private Charset getCharset() { diff --git a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java index c1f0e17bf7..a2aba3addd 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/protobuf/ProtobufDecoder.java @@ -36,6 +36,7 @@ import org.springframework.core.ResolvableType; import org.springframework.core.codec.Decoder; import org.springframework.core.codec.DecodingException; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.lang.Nullable; import org.springframework.util.Assert; @@ -101,10 +102,24 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements DecoderBy default in 5.1 this is set to 64K. In 5.2 the default for this limit + * is set to 256K. + * @param maxMessageSize the max size per message, or -1 for unlimited + */ public void setMaxMessageSize(int maxMessageSize) { this.maxMessageSize = maxMessageSize; } + /** + * Return the {@link #setMaxMessageSize configured} message size limit. + * @since 5.1.11 + */ + public int getMaxMessageSize() { + return this.maxMessageSize; + } + @Override public boolean canDecode(ResolvableType elementType, @Nullable MimeType mimeType) { @@ -127,7 +142,7 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements Decoder decodeToMono(Publisher inputStream, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { - return DataBufferUtils.join(inputStream).map(dataBuffer -> { + return DataBufferUtils.join(inputStream, getMaxMessageSize()).map(dataBuffer -> { try { Message.Builder builder = getMessageBuilder(elementType.toClass()); ByteBuffer buffer = dataBuffer.asByteBuffer(); @@ -198,9 +213,9 @@ public class ProtobufDecoder extends ProtobufCodecSupport implements Decoder this.maxMessageSize) { - throw new DecodingException( - "The number of bytes to read from the incoming stream " + + if (this.maxMessageSize > 0 && this.messageBytesToRead > this.maxMessageSize) { + throw new DataBufferLimitException( + "The number of bytes to read for message " + "(" + this.messageBytesToRead + ") exceeds " + "the configured limit (" + this.maxMessageSize + ")"); } diff --git a/spring-web/src/main/java/org/springframework/http/codec/support/BaseDefaultCodecs.java b/spring-web/src/main/java/org/springframework/http/codec/support/BaseDefaultCodecs.java index 055a0ed29b..f3034ad935 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/support/BaseDefaultCodecs.java +++ b/spring-web/src/main/java/org/springframework/http/codec/support/BaseDefaultCodecs.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -20,6 +20,7 @@ import java.util.ArrayList; import java.util.Collections; import java.util.List; +import org.springframework.core.codec.AbstractDataBufferDecoder; import org.springframework.core.codec.ByteArrayDecoder; import org.springframework.core.codec.ByteArrayEncoder; import org.springframework.core.codec.ByteBufferDecoder; @@ -38,6 +39,7 @@ import org.springframework.http.codec.FormHttpMessageReader; import org.springframework.http.codec.HttpMessageReader; import org.springframework.http.codec.HttpMessageWriter; import org.springframework.http.codec.ResourceHttpMessageWriter; +import org.springframework.http.codec.json.AbstractJackson2Decoder; import org.springframework.http.codec.json.Jackson2JsonDecoder; import org.springframework.http.codec.json.Jackson2JsonEncoder; import org.springframework.http.codec.json.Jackson2SmileDecoder; @@ -95,6 +97,9 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { @Nullable private Encoder jaxb2Encoder; + @Nullable + private Integer maxInMemorySize; + private boolean enableLoggingRequestDetails = false; private boolean registerDefaults = true; @@ -130,6 +135,16 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { this.jaxb2Encoder = encoder; } + @Override + public void maxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + @Nullable + protected Integer maxInMemorySize() { + return this.maxInMemorySize; + } + @Override public void enableLoggingRequestDetails(boolean enable) { this.enableLoggingRequestDetails = enable; @@ -155,17 +170,20 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { return Collections.emptyList(); } List> readers = new ArrayList<>(); - readers.add(new DecoderHttpMessageReader<>(new ByteArrayDecoder())); - readers.add(new DecoderHttpMessageReader<>(new ByteBufferDecoder())); - readers.add(new DecoderHttpMessageReader<>(new DataBufferDecoder())); - readers.add(new DecoderHttpMessageReader<>(new ResourceDecoder())); - readers.add(new DecoderHttpMessageReader<>(StringDecoder.textPlainOnly())); + readers.add(new DecoderHttpMessageReader<>(init(new ByteArrayDecoder()))); + readers.add(new DecoderHttpMessageReader<>(init(new ByteBufferDecoder()))); + readers.add(new DecoderHttpMessageReader<>(init(new DataBufferDecoder()))); + readers.add(new DecoderHttpMessageReader<>(init(new ResourceDecoder()))); + readers.add(new DecoderHttpMessageReader<>(init(StringDecoder.textPlainOnly()))); if (protobufPresent) { - Decoder decoder = this.protobufDecoder != null ? this.protobufDecoder : new ProtobufDecoder(); + Decoder decoder = this.protobufDecoder != null ? this.protobufDecoder : init(new ProtobufDecoder()); readers.add(new DecoderHttpMessageReader<>(decoder)); } FormHttpMessageReader formReader = new FormHttpMessageReader(); + if (this.maxInMemorySize != null) { + formReader.setMaxInMemorySize(this.maxInMemorySize); + } formReader.setEnableLoggingRequestDetails(this.enableLoggingRequestDetails); readers.add(formReader); @@ -174,6 +192,28 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { return readers; } + private > T init(T decoder) { + if (this.maxInMemorySize != null) { + if (decoder instanceof AbstractDataBufferDecoder) { + ((AbstractDataBufferDecoder) decoder).setMaxInMemorySize(this.maxInMemorySize); + } + if (decoder instanceof ProtobufDecoder) { + ((ProtobufDecoder) decoder).setMaxMessageSize(this.maxInMemorySize); + } + if (jackson2Present) { + if (decoder instanceof AbstractJackson2Decoder) { + ((AbstractJackson2Decoder) decoder).setMaxInMemorySize(this.maxInMemorySize); + } + } + if (jaxb2Present) { + if (decoder instanceof Jaxb2XmlDecoder) { + ((Jaxb2XmlDecoder) decoder).setMaxInMemorySize(this.maxInMemorySize); + } + } + } + return decoder; + } + /** * Hook for client or server specific typed readers. */ @@ -189,13 +229,13 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { } List> readers = new ArrayList<>(); if (jackson2Present) { - readers.add(new DecoderHttpMessageReader<>(getJackson2JsonDecoder())); + readers.add(new DecoderHttpMessageReader<>(init(getJackson2JsonDecoder()))); } if (jackson2SmilePresent) { - readers.add(new DecoderHttpMessageReader<>(new Jackson2SmileDecoder())); + readers.add(new DecoderHttpMessageReader<>(init(new Jackson2SmileDecoder()))); } if (jaxb2Present) { - Decoder decoder = this.jaxb2Decoder != null ? this.jaxb2Decoder : new Jaxb2XmlDecoder(); + Decoder decoder = this.jaxb2Decoder != null ? this.jaxb2Decoder : init(new Jaxb2XmlDecoder()); readers.add(new DecoderHttpMessageReader<>(decoder)); } extendObjectReaders(readers); @@ -216,7 +256,7 @@ class BaseDefaultCodecs implements CodecConfigurer.DefaultCodecs { return Collections.emptyList(); } List> result = new ArrayList<>(); - result.add(new DecoderHttpMessageReader<>(StringDecoder.allMimeTypes())); + result.add(new DecoderHttpMessageReader<>(init(StringDecoder.allMimeTypes()))); return result; } diff --git a/spring-web/src/main/java/org/springframework/http/codec/support/ServerDefaultCodecsImpl.java b/spring-web/src/main/java/org/springframework/http/codec/support/ServerDefaultCodecsImpl.java index 15461d11f4..37e924cd7e 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/support/ServerDefaultCodecsImpl.java +++ b/spring-web/src/main/java/org/springframework/http/codec/support/ServerDefaultCodecsImpl.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -39,10 +39,18 @@ class ServerDefaultCodecsImpl extends BaseDefaultCodecs implements ServerCodecCo DefaultServerCodecConfigurer.class.getClassLoader()); + @Nullable + private HttpMessageReader multipartReader; + @Nullable private Encoder sseEncoder; + @Override + public void multipartReader(HttpMessageReader reader) { + this.multipartReader = reader; + } + @Override public void serverSentEventEncoder(Encoder encoder) { this.sseEncoder = encoder; @@ -51,10 +59,18 @@ class ServerDefaultCodecsImpl extends BaseDefaultCodecs implements ServerCodecCo @Override protected void extendTypedReaders(List> typedReaders) { + if (this.multipartReader != null) { + typedReaders.add(this.multipartReader); + return; + } if (synchronossMultipartPresent) { boolean enable = isEnableLoggingRequestDetails(); SynchronossPartHttpMessageReader partReader = new SynchronossPartHttpMessageReader(); + Integer size = maxInMemorySize(); + if (size != null) { + partReader.setMaxInMemorySize(size); + } partReader.setEnableLoggingRequestDetails(enable); typedReaders.add(partReader); diff --git a/spring-web/src/main/java/org/springframework/http/codec/xml/Jaxb2XmlDecoder.java b/spring-web/src/main/java/org/springframework/http/codec/xml/Jaxb2XmlDecoder.java index 1b87b0e6d5..7fd886fb0d 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/xml/Jaxb2XmlDecoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/xml/Jaxb2XmlDecoder.java @@ -43,6 +43,7 @@ import org.springframework.core.codec.CodecException; import org.springframework.core.codec.DecodingException; import org.springframework.core.codec.Hints; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.log.LogFormatUtils; import org.springframework.lang.Nullable; import org.springframework.util.Assert; @@ -78,6 +79,8 @@ public class Jaxb2XmlDecoder extends AbstractDecoder { private Function unmarshallerProcessor = Function.identity(); + private int maxInMemorySize = -1; + public Jaxb2XmlDecoder() { super(MimeTypeUtils.APPLICATION_XML, MimeTypeUtils.TEXT_XML); @@ -110,6 +113,29 @@ public class Jaxb2XmlDecoder extends AbstractDecoder { return this.unmarshallerProcessor; } + /** + * Set the max number of bytes that can be buffered by this decoder. + * This is either the size of the entire input when decoding as a whole, or when + * using async parsing with Aalto XML, it is the size of one top-level XML tree. + * When the limit is exceeded, {@link DataBufferLimitException} is raised. + *

By default in 5.1 this is set to -1, unlimited. In 5.2 the default + * value for this limit is set to 256K. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + this.xmlEventDecoder.setMaxInMemorySize(byteCount); + } + + /** + * Return the {@link #setMaxInMemorySize configured} byte count limit. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + @Override public boolean canDecode(ResolvableType elementType, @Nullable MimeType mimeType) { diff --git a/spring-web/src/main/java/org/springframework/http/codec/xml/XmlEventDecoder.java b/spring-web/src/main/java/org/springframework/http/codec/xml/XmlEventDecoder.java index 5f1665399f..46d3c4e139 100644 --- a/spring-web/src/main/java/org/springframework/http/codec/xml/XmlEventDecoder.java +++ b/spring-web/src/main/java/org/springframework/http/codec/xml/XmlEventDecoder.java @@ -40,6 +40,7 @@ import reactor.core.publisher.Flux; import org.springframework.core.ResolvableType; import org.springframework.core.codec.AbstractDecoder; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.lang.Nullable; import org.springframework.util.ClassUtils; @@ -88,26 +89,51 @@ public class XmlEventDecoder extends AbstractDecoder { boolean useAalto = aaltoPresent; + private int maxInMemorySize = -1; + public XmlEventDecoder() { super(MimeTypeUtils.APPLICATION_XML, MimeTypeUtils.TEXT_XML); } + /** + * Set the max number of bytes that can be buffered by this decoder. This + * is either the size the entire input when decoding as a whole, or when + * using async parsing via Aalto XML, it is size one top-level XML tree. + * When the limit is exceeded, {@link DataBufferLimitException} is raised. + *

By default in 5.1 this is set to -1, unlimited. In 5.2 the default + * value for this limit is set to 256K. + * @param byteCount the max number of bytes to buffer, or -1 for unlimited + * @since 5.1.11 + */ + public void setMaxInMemorySize(int byteCount) { + this.maxInMemorySize = byteCount; + } + + /** + * Return the {@link #setMaxInMemorySize configured} byte count limit. + * @since 5.1.11 + */ + public int getMaxInMemorySize() { + return this.maxInMemorySize; + } + + @Override @SuppressWarnings({"rawtypes", "unchecked"}) // on JDK 9 where XMLEventReader is Iterator public Flux decode(Publisher input, ResolvableType elementType, @Nullable MimeType mimeType, @Nullable Map hints) { if (this.useAalto) { - AaltoDataBufferToXmlEvent mapper = new AaltoDataBufferToXmlEvent(); + AaltoDataBufferToXmlEvent mapper = new AaltoDataBufferToXmlEvent(this.maxInMemorySize); return Flux.from(input) .flatMapIterable(mapper) .doFinally(signalType -> mapper.endOfInput()); } else { - return DataBufferUtils.join(input). - flatMapIterable(buffer -> { + return DataBufferUtils.join(input, getMaxInMemorySize()) + .flatMapIterable(buffer -> { try { InputStream is = buffer.asInputStream(); Iterator eventReader = inputFactory.createXMLEventReader(is); @@ -139,10 +165,22 @@ public class XmlEventDecoder extends AbstractDecoder { private final XMLEventAllocator eventAllocator = EventAllocatorImpl.getDefaultInstance(); + private final int maxInMemorySize; + + private int byteCount; + + private int elementDepth; + + + public AaltoDataBufferToXmlEvent(int maxInMemorySize) { + this.maxInMemorySize = maxInMemorySize; + } + @Override public List apply(DataBuffer dataBuffer) { try { + increaseByteCount(dataBuffer); this.streamReader.getInputFeeder().feedInput(dataBuffer.asByteBuffer()); List events = new ArrayList<>(); while (true) { @@ -156,8 +194,12 @@ public class XmlEventDecoder extends AbstractDecoder { if (event.isEndDocument()) { break; } + checkDepthAndResetByteCount(event); } } + if (this.maxInMemorySize > 0 && this.byteCount > this.maxInMemorySize) { + raiseLimitException(); + } return events; } catch (XMLStreamException ex) { @@ -168,6 +210,35 @@ public class XmlEventDecoder extends AbstractDecoder { } } + private void increaseByteCount(DataBuffer dataBuffer) { + if (this.maxInMemorySize > 0) { + if (dataBuffer.readableByteCount() > Integer.MAX_VALUE - this.byteCount) { + raiseLimitException(); + } + else { + this.byteCount += dataBuffer.readableByteCount(); + } + } + } + + private void checkDepthAndResetByteCount(XMLEvent event) { + if (this.maxInMemorySize > 0) { + if (event.isStartElement()) { + this.byteCount = this.elementDepth == 1 ? 0 : this.byteCount; + this.elementDepth++; + } + else if (event.isEndElement()) { + this.elementDepth--; + this.byteCount = this.elementDepth == 1 ? 0 : this.byteCount; + } + } + } + + private void raiseLimitException() { + throw new DataBufferLimitException( + "Exceeded limit on max bytes per XML top-level node: " + this.maxInMemorySize); + } + public void endOfInput() { this.streamReader.getInputFeeder().endOfInput(); } diff --git a/spring-web/src/test/java/org/springframework/http/codec/json/Jackson2TokenizerTests.java b/spring-web/src/test/java/org/springframework/http/codec/json/Jackson2TokenizerTests.java index 6e9ef25569..5c08550c07 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/json/Jackson2TokenizerTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/json/Jackson2TokenizerTests.java @@ -20,7 +20,6 @@ import java.io.IOException; import java.io.UncheckedIOException; import java.nio.charset.StandardCharsets; import java.util.List; -import java.util.function.Consumer; import com.fasterxml.jackson.core.JsonFactory; import com.fasterxml.jackson.core.TreeNode; @@ -36,6 +35,7 @@ import reactor.test.StepVerifier; import org.springframework.core.codec.DecodingException; import org.springframework.core.io.buffer.AbstractLeakCheckingTestCase; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; import static java.util.Arrays.*; import static java.util.Collections.*; @@ -181,11 +181,68 @@ public class Jackson2TokenizerTests extends AbstractLeakCheckingTestCase { testTokenize(asList("[1", ",2,", "3]"), asList("1", "2", "3"), true); } + private void testTokenize(List input, List output, boolean tokenize) { + StepVerifier.FirstStep builder = StepVerifier.create(decode(input, tokenize, -1)); + output.forEach(expected -> builder.assertNext(actual -> { + try { + JSONAssert.assertEquals(expected, actual, true); + } + catch (JSONException ex) { + throw new RuntimeException(ex); + } + })); + builder.verifyComplete(); + } + + @Test + public void testLimit() { + + List source = asList("[", + "{", "\"id\":1,\"name\":\"Dan\"", "},", + "{", "\"id\":2,\"name\":\"Ron\"", "},", + "{", "\"id\":3,\"name\":\"Bartholomew\"", "}", + "]"); + + String expected = String.join("", source); + int maxInMemorySize = expected.length(); + + StepVerifier.create(decode(source, false, maxInMemorySize)) + .expectNext(expected) + .verifyComplete(); + + StepVerifier.create(decode(source, false, maxInMemorySize - 1)) + .expectError(DataBufferLimitException.class); + } + + @Test + public void testLimitTokenized() { + + List source = asList("[", + "{", "\"id\":1, \"name\":\"Dan\"", "},", + "{", "\"id\":2, \"name\":\"Ron\"", "},", + "{", "\"id\":3, \"name\":\"Bartholomew\"", "}", + "]"); + + String expected = "{\"id\":3,\"name\":\"Bartholomew\"}"; + int maxInMemorySize = expected.length(); + + StepVerifier.create(decode(source, true, maxInMemorySize)) + .expectNext("{\"id\":1,\"name\":\"Dan\"}") + .expectNext("{\"id\":2,\"name\":\"Ron\"}") + .expectNext(expected) + .verifyComplete(); + + StepVerifier.create(decode(source, true, maxInMemorySize - 1)) + .expectNext("{\"id\":1,\"name\":\"Dan\"}") + .expectNext("{\"id\":2,\"name\":\"Ron\"}") + .verifyError(DataBufferLimitException.class); + } + @Test public void errorInStream() { DataBuffer buffer = stringBuffer("{\"id\":1,\"name\":"); Flux source = Flux.just(buffer).concatWith(Flux.error(new RuntimeException())); - Flux result = Jackson2Tokenizer.tokenize(source, this.jsonFactory, this.objectMapper, true); + Flux result = Jackson2Tokenizer.tokenize(source, this.jsonFactory, this.objectMapper, true, -1); StepVerifier.create(result) .expectError(RuntimeException.class) @@ -195,7 +252,7 @@ public class Jackson2TokenizerTests extends AbstractLeakCheckingTestCase { @Test // SPR-16521 public void jsonEOFExceptionIsWrappedAsDecodingError() { Flux source = Flux.just(stringBuffer("{\"status\": \"noClosingQuote}")); - Flux tokens = Jackson2Tokenizer.tokenize(source, this.jsonFactory, this.objectMapper, false); + Flux tokens = Jackson2Tokenizer.tokenize(source, this.jsonFactory, this.objectMapper, false, -1); StepVerifier.create(tokens) .expectError(DecodingException.class) @@ -203,12 +260,13 @@ public class Jackson2TokenizerTests extends AbstractLeakCheckingTestCase { } - private void testTokenize(List source, List expected, boolean tokenizeArrayElements) { + private Flux decode(List source, boolean tokenize, int maxInMemorySize) { + Flux tokens = Jackson2Tokenizer.tokenize( Flux.fromIterable(source).map(this::stringBuffer), - this.jsonFactory, this.objectMapper, tokenizeArrayElements); + this.jsonFactory, this.objectMapper, tokenize, maxInMemorySize); - Flux result = tokens + return tokens .map(tokenBuffer -> { try { TreeNode root = this.objectMapper.readTree(tokenBuffer.asParser()); @@ -218,10 +276,6 @@ public class Jackson2TokenizerTests extends AbstractLeakCheckingTestCase { throw new UncheckedIOException(ex); } }); - - StepVerifier.FirstStep builder = StepVerifier.create(result); - expected.forEach(s -> builder.assertNext(new JSONAssertConsumer(s))); - builder.verifyComplete(); } private DataBuffer stringBuffer(String value) { @@ -231,24 +285,4 @@ public class Jackson2TokenizerTests extends AbstractLeakCheckingTestCase { return buffer; } - - private static class JSONAssertConsumer implements Consumer { - - private final String expected; - - JSONAssertConsumer(String expected) { - this.expected = expected; - } - - @Override - public void accept(String s) { - try { - JSONAssert.assertEquals(this.expected, s, true); - } - catch (JSONException ex) { - throw new RuntimeException(ex); - } - } - } - } diff --git a/spring-web/src/test/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReaderTests.java b/spring-web/src/test/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReaderTests.java index d5052aa24e..74d3fb0db0 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReaderTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/multipart/SynchronossPartHttpMessageReaderTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -17,15 +17,20 @@ package org.springframework.http.codec.multipart; import java.io.File; +import java.io.IOException; import java.time.Duration; import java.util.Map; +import java.util.function.Consumer; import org.junit.Test; +import org.reactivestreams.Subscription; +import reactor.core.publisher.BaseSubscriber; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; import reactor.test.StepVerifier; import org.springframework.core.ResolvableType; +import org.springframework.core.codec.DecodingException; import org.springframework.core.io.ClassPathResource; import org.springframework.core.io.buffer.DataBuffer; import org.springframework.core.io.buffer.DataBufferUtils; @@ -38,23 +43,31 @@ import org.springframework.mock.http.client.reactive.test.MockClientHttpRequest; import org.springframework.mock.http.server.reactive.test.MockServerHttpRequest; import org.springframework.util.MultiValueMap; -import static java.util.Collections.*; -import static org.junit.Assert.*; -import static org.springframework.core.ResolvableType.*; -import static org.springframework.http.HttpHeaders.*; -import static org.springframework.http.MediaType.*; +import static java.util.Collections.emptyMap; +import static org.hamcrest.core.StringStartsWith.startsWith; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertNotNull; +import static org.junit.Assert.assertThat; +import static org.junit.Assert.assertTrue; +import static org.springframework.core.ResolvableType.forClassWithGenerics; +import static org.springframework.http.HttpHeaders.CONTENT_TYPE; +import static org.springframework.http.MediaType.MULTIPART_FORM_DATA; /** * Unit tests for {@link SynchronossPartHttpMessageReader}. * * @author Sebastien Deleuze * @author Rossen Stoyanchev + * @author Brian Clozel */ public class SynchronossPartHttpMessageReaderTests { private final MultipartHttpMessageReader reader = new MultipartHttpMessageReader(new SynchronossPartHttpMessageReader()); + private static final ResolvableType PARTS_ELEMENT_TYPE = + forClassWithGenerics(MultiValueMap.class, String.class, Part.class); @Test public void canRead() { @@ -86,10 +99,10 @@ public class SynchronossPartHttpMessageReaderTests { MultiValueMap parts = this.reader.readMono(elementType, request, emptyMap()).block(); assertEquals(2, parts.size()); - assertTrue(parts.containsKey("fooPart")); - Part part = parts.getFirst("fooPart"); + assertTrue(parts.containsKey("filePart")); + Part part = parts.getFirst("filePart"); assertTrue(part instanceof FilePart); - assertEquals("fooPart", part.name()); + assertEquals("filePart", part.name()); assertEquals("foo.txt", ((FilePart) part).filename()); DataBuffer buffer = DataBufferUtils.join(part.content()).block(); assertEquals(12, buffer.readableByteCount()); @@ -97,24 +110,23 @@ public class SynchronossPartHttpMessageReaderTests { buffer.read(byteContent); assertEquals("Lorem Ipsum.", new String(byteContent)); - assertTrue(parts.containsKey("barPart")); - part = parts.getFirst("barPart"); + assertTrue(parts.containsKey("textPart")); + part = parts.getFirst("textPart"); assertTrue(part instanceof FormFieldPart); - assertEquals("barPart", part.name()); - assertEquals("bar", ((FormFieldPart) part).value()); + assertEquals("textPart", part.name()); + assertEquals("sample-text", ((FormFieldPart) part).value()); } @Test // SPR-16545 - public void transferTo() { + public void transferTo() throws IOException { ServerHttpRequest request = generateMultipartRequest(); - ResolvableType elementType = forClassWithGenerics(MultiValueMap.class, String.class, Part.class); - MultiValueMap parts = this.reader.readMono(elementType, request, emptyMap()).block(); + MultiValueMap parts = this.reader.readMono(PARTS_ELEMENT_TYPE, request, emptyMap()).block(); assertNotNull(parts); - FilePart part = (FilePart) parts.getFirst("fooPart"); + FilePart part = (FilePart) parts.getFirst("filePart"); assertNotNull(part); - File dest = new File(System.getProperty("java.io.tmpdir") + "/" + part.filename()); + File dest = File.createTempFile(part.filename(), "multipart"); part.transferTo(dest).block(Duration.ofSeconds(5)); assertTrue(dest.exists()); @@ -125,22 +137,65 @@ public class SynchronossPartHttpMessageReaderTests { @Test public void bodyError() { ServerHttpRequest request = generateErrorMultipartRequest(); - ResolvableType elementType = forClassWithGenerics(MultiValueMap.class, String.class, Part.class); - StepVerifier.create(this.reader.readMono(elementType, request, emptyMap())).verifyError(); + StepVerifier.create(this.reader.readMono(PARTS_ELEMENT_TYPE, request, emptyMap())).verifyError(); } + @Test + public void readPartsWithoutDemand() { + ServerHttpRequest request = generateMultipartRequest(); + Mono> parts = this.reader.readMono(PARTS_ELEMENT_TYPE, request, emptyMap()); + ZeroDemandSubscriber subscriber = new ZeroDemandSubscriber(); + parts.subscribe(subscriber); + subscriber.cancel(); + } + + @Test + public void readTooManyParts() { + testMultipartExceptions(reader -> reader.setMaxParts(1), ex -> { + assertEquals(DecodingException.class, ex.getClass()); + assertThat(ex.getMessage(), startsWith("Failure while parsing part[2]")); + assertEquals("Too many parts (2 allowed)", ex.getCause().getMessage()); + }); + } + + @Test + public void readFilePartTooBig() { + testMultipartExceptions(reader -> reader.setMaxDiskUsagePerPart(5), ex -> { + assertEquals(DecodingException.class, ex.getClass()); + assertThat(ex.getMessage(), startsWith("Failure while parsing part[1]")); + assertEquals("Part[1] exceeded the disk usage limit of 5 bytes", ex.getCause().getMessage()); + }); + } + + @Test + public void readPartHeadersTooBig() { + testMultipartExceptions(reader -> reader.setMaxInMemorySize(1), ex -> { + assertEquals(DecodingException.class, ex.getClass()); + assertThat(ex.getMessage(), startsWith("Failure while parsing part[1]")); + assertEquals("Part[1] exceeded the in-memory limit of 1 bytes", ex.getCause().getMessage()); + }); + } + + private void testMultipartExceptions( + Consumer configurer, Consumer assertions) { + + SynchronossPartHttpMessageReader reader = new SynchronossPartHttpMessageReader(); + configurer.accept(reader); + MultipartHttpMessageReader multipartReader = new MultipartHttpMessageReader(reader); + StepVerifier.create(multipartReader.readMono(PARTS_ELEMENT_TYPE, generateMultipartRequest(), emptyMap())) + .consumeErrorWith(assertions) + .verify(); + } private ServerHttpRequest generateMultipartRequest() { - MultipartBodyBuilder partsBuilder = new MultipartBodyBuilder(); - partsBuilder.part("fooPart", new ClassPathResource("org/springframework/http/codec/multipart/foo.txt")); - partsBuilder.part("barPart", "bar"); + partsBuilder.part("filePart", new ClassPathResource("org/springframework/http/codec/multipart/foo.txt")); + partsBuilder.part("textPart", "sample-text"); MockClientHttpRequest outputMessage = new MockClientHttpRequest(HttpMethod.POST, "/"); new MultipartHttpMessageWriter() .write(Mono.just(partsBuilder.build()), null, MediaType.MULTIPART_FORM_DATA, outputMessage, null) .block(Duration.ofSeconds(5)); - return MockServerHttpRequest.post("/") .contentType(outputMessage.getHeaders().getContentType()) .body(outputMessage.getBody()); @@ -152,4 +207,12 @@ public class SynchronossPartHttpMessageReaderTests { .body(Flux.just(new DefaultDataBufferFactory().wrap("invalid content".getBytes()))); } + private static class ZeroDemandSubscriber extends BaseSubscriber> { + + @Override + protected void hookOnSubscribe(Subscription subscription) { + // Just subscribe without requesting + } + } + } diff --git a/spring-web/src/test/java/org/springframework/http/codec/support/ServerCodecConfigurerTests.java b/spring-web/src/test/java/org/springframework/http/codec/support/ServerCodecConfigurerTests.java index 4416507571..b36cdd0ca7 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/support/ServerCodecConfigurerTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/support/ServerCodecConfigurerTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -124,13 +124,41 @@ public class ServerCodecConfigurerTests { .filter(e -> e == encoder).orElse(null)); } + @Test + public void maxInMemorySize() { + int size = 99; + this.configurer.defaultCodecs().maxInMemorySize(size); + List> readers = this.configurer.getReaders(); + assertEquals(13, readers.size()); + assertEquals(size, ((ByteArrayDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((ByteBufferDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((DataBufferDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((ResourceDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((StringDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((ProtobufDecoder) getNextDecoder(readers)).getMaxMessageSize()); + assertEquals(size, ((FormHttpMessageReader) nextReader(readers)).getMaxInMemorySize()); + assertEquals(size, ((SynchronossPartHttpMessageReader) nextReader(readers)).getMaxInMemorySize()); + + MultipartHttpMessageReader multipartReader = (MultipartHttpMessageReader) nextReader(readers); + SynchronossPartHttpMessageReader reader = (SynchronossPartHttpMessageReader) multipartReader.getPartReader(); + assertEquals(size, (reader).getMaxInMemorySize()); + + assertEquals(size, ((Jackson2JsonDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((Jackson2SmileDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((Jaxb2XmlDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + assertEquals(size, ((StringDecoder) getNextDecoder(readers)).getMaxInMemorySize()); + } private Decoder getNextDecoder(List> readers) { - HttpMessageReader reader = readers.get(this.index.getAndIncrement()); + HttpMessageReader reader = nextReader(readers); assertEquals(DecoderHttpMessageReader.class, reader.getClass()); return ((DecoderHttpMessageReader) reader).getDecoder(); } + private HttpMessageReader nextReader(List> readers) { + return readers.get(this.index.getAndIncrement()); + } + private Encoder getNextEncoder(List> writers) { HttpMessageWriter writer = writers.get(this.index.getAndIncrement()); assertEquals(EncoderHttpMessageWriter.class, writer.getClass()); diff --git a/spring-web/src/test/java/org/springframework/http/codec/xml/XmlEventDecoderTests.java b/spring-web/src/test/java/org/springframework/http/codec/xml/XmlEventDecoderTests.java index 8babb0b614..5e4c394145 100644 --- a/spring-web/src/test/java/org/springframework/http/codec/xml/XmlEventDecoderTests.java +++ b/spring-web/src/test/java/org/springframework/http/codec/xml/XmlEventDecoderTests.java @@ -1,5 +1,5 @@ /* - * Copyright 2002-2018 the original author or authors. + * Copyright 2002-2019 the original author or authors. * * Licensed under the Apache License, Version 2.0 (the "License"); * you may not use this file except in compliance with the License. @@ -28,8 +28,10 @@ import reactor.test.StepVerifier; import org.springframework.core.io.buffer.AbstractLeakCheckingTestCase; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferLimitException; -import static org.junit.Assert.*; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertTrue; /** * @author Arjen Poutsma @@ -44,11 +46,12 @@ public class XmlEventDecoderTests extends AbstractLeakCheckingTestCase { private XmlEventDecoder decoder = new XmlEventDecoder(); + @Test public void toXMLEventsAalto() { Flux events = - this.decoder.decode(stringBuffer(XML), null, null, Collections.emptyMap()); + this.decoder.decode(stringBufferMono(XML), null, null, Collections.emptyMap()); StepVerifier.create(events) .consumeNextWith(e -> assertTrue(e.isStartDocument())) @@ -69,7 +72,7 @@ public class XmlEventDecoderTests extends AbstractLeakCheckingTestCase { decoder.useAalto = false; Flux events = - this.decoder.decode(stringBuffer(XML), null, null, Collections.emptyMap()); + this.decoder.decode(stringBufferMono(XML), null, null, Collections.emptyMap()); StepVerifier.create(events) .consumeNextWith(e -> assertTrue(e.isStartDocument())) @@ -86,10 +89,32 @@ public class XmlEventDecoderTests extends AbstractLeakCheckingTestCase { .verify(); } + @Test + public void toXMLEventsWithLimit() { + + this.decoder.setMaxInMemorySize(6); + + Flux source = Flux.just( + "", "", "foofoo", "", "", "barbarbar", "", ""); + + Flux events = this.decoder.decode( + source.map(this::stringBuffer), null, null, Collections.emptyMap()); + + StepVerifier.create(events) + .consumeNextWith(e -> assertTrue(e.isStartDocument())) + .consumeNextWith(e -> assertStartElement(e, "pojo")) + .consumeNextWith(e -> assertStartElement(e, "foo")) + .consumeNextWith(e -> assertCharacters(e, "foofoo")) + .consumeNextWith(e -> assertEndElement(e, "foo")) + .consumeNextWith(e -> assertStartElement(e, "bar")) + .expectError(DataBufferLimitException.class) + .verify(); + } + @Test public void decodeErrorAalto() { Flux source = Flux.concat( - stringBuffer(""), + stringBufferMono(""), Flux.error(new RuntimeException())); Flux events = @@ -107,7 +132,7 @@ public class XmlEventDecoderTests extends AbstractLeakCheckingTestCase { decoder.useAalto = false; Flux source = Flux.concat( - stringBuffer(""), + stringBufferMono(""), Flux.error(new RuntimeException())); Flux events = @@ -133,13 +158,15 @@ public class XmlEventDecoderTests extends AbstractLeakCheckingTestCase { assertEquals(expectedData, event.asCharacters().getData()); } - private Mono stringBuffer(String value) { - return Mono.defer(() -> { - byte[] bytes = value.getBytes(StandardCharsets.UTF_8); - DataBuffer buffer = this.bufferFactory.allocateBuffer(bytes.length); - buffer.write(bytes); - return Mono.just(buffer); - }); + private DataBuffer stringBuffer(String value) { + byte[] bytes = value.getBytes(StandardCharsets.UTF_8); + DataBuffer buffer = this.bufferFactory.allocateBuffer(bytes.length); + buffer.write(bytes); + return buffer; + } + + private Mono stringBufferMono(String value) { + return Mono.defer(() -> Mono.just(stringBuffer(value))); } } diff --git a/src/docs/asciidoc/web/webflux.adoc b/src/docs/asciidoc/web/webflux.adoc index 330ef75635..5789f29527 100644 --- a/src/docs/asciidoc/web/webflux.adoc +++ b/src/docs/asciidoc/web/webflux.adoc @@ -761,6 +761,33 @@ for repeated, map-like access to parts, or otherwise rely on the `SynchronossPartHttpMessageReader` for a one-time access to `Flux`. +[[webflux-codecs-limits]] +==== Limits + +`Decoder` and `HttpMessageReader` implementations that buffer some or all of the input +stream can be configured with a limit on the maximum number of bytes to buffer in memory. +In some cases buffering occurs because input is aggregated and represented as a single +object, e.g. controller method with `@RequestBody byte[]`, `x-www-form-urlencoded` data, +and so on. Buffering can also occurs with streaming, when splitting the input stream, +e.g. delimited text, a stream of JSON objects, and so on. For those streaming cases, the +limit applies to the number of bytes associted with one object in the stream. + +To configure buffer sizes, you can check if a given `Decoder` or `HttpMessageReader` +exposes a `maxInMemorySize` property and if so the Javadoc will have details about default +values. In WebFlux, the `ServerCodecConfigurer` provides a +<> from where to set all codecs, through the +`maxInMemorySize` property for default codecs. + +For <> the `maxInMemorySize` property limits +the size of non-file parts. For file parts it determines the threshold at which the part +is written to disk. For file parts written to disk, there is an additional +`maxDiskUsagePerPart` property to limit the amount of disk space per part. There is also +a `maxParts` property to limit the overall number of parts in a multipart request. +To configure all 3 in WebFlux, you'll need to supply a pre-configured instance of +`MultipartHttpMessageReader` to `ServerCodecConfigurer`. + + + [[webflux-codecs-streaming]] ==== Streaming [.small]#<>#