diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/handler/predicate/ReadBodyPredicateFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/handler/predicate/ReadBodyPredicateFactory.java index f4d32198..6ead30a2 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/handler/predicate/ReadBodyPredicateFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/handler/predicate/ReadBodyPredicateFactory.java @@ -18,10 +18,17 @@ package org.springframework.cloud.gateway.handler.predicate; import java.util.Collections; +import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Set; +import java.util.function.BiConsumer; +import java.util.function.BinaryOperator; import java.util.function.Function; import java.util.function.Predicate; +import java.util.function.Supplier; +import java.util.stream.Collector; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; @@ -37,7 +44,6 @@ import org.springframework.core.io.buffer.DataBuffer; import org.springframework.http.HttpHeaders; import org.springframework.http.ReactiveHttpInputMessage; import org.springframework.http.codec.HttpMessageReader; -import org.springframework.http.codec.ServerCodecConfigurer; import org.springframework.web.reactive.function.BodyInserter; import org.springframework.web.reactive.function.BodyInserters; @@ -45,6 +51,8 @@ import org.springframework.web.reactive.function.server.HandlerStrategies; import org.springframework.web.server.ServerWebExchange; import static org.springframework.cloud.gateway.filter.AdaptCachedBodyGlobalFilter.CACHED_REQUEST_BODY_KEY; +import static org.springframework.cloud.gateway.handler.predicate.ReadBodyPredicateFactory.DataBufferMapCollector.BODY_ONE; +import static org.springframework.cloud.gateway.handler.predicate.ReadBodyPredicateFactory.DataBufferMapCollector.BODY_TWO; /** * This predicate is BETA and may be subject to change in a future release. @@ -73,55 +81,101 @@ public class ReadBodyPredicateFactory // exception will be thrown. The below if/else caches the body object as a request attribute in the ServerWebExchange // so if this filter is run more than once (due to more than one route using it) we do not try to read the // request body multiple times - if(cachedBody != null) { + if (cachedBody != null) { try { boolean test = config.predicate.test(cachedBody); exchange.getAttributes().put(TEST_ATTRIBUTE, test); return Mono.just(test); - } catch(ClassCastException e) { - if(LOGGER.isDebugEnabled()) { + } catch (ClassCastException e) { + if (LOGGER.isDebugEnabled()) { LOGGER.debug("Predicate test failed because class in predicate does not match the cached body object", e); } } return Mono.just(false); } else { - Flux origBody = exchange.getRequest().getBody().flatMap(dataBuffer -> { - ResolvableType type =ResolvableType.forClass(inClass); - for(HttpMessageReader messageReader: messageReaders) { - if(messageReader.canRead(type, exchange.getRequest().getHeaders().getContentType())) { - ReactiveHttpInputMessage inputMessage = new ReadBodyReactiveHttpInputMessage( - Flux.just(dataBuffer.factory().allocateBuffer().write(dataBuffer.asByteBuffer())), - exchange.getRequest().getHeaders()); - Function mapper = (bodyObj) -> { - exchange.getAttributes().put(CACHE_REQUEST_BODY_OBJECT_KEY, bodyObj); - boolean test = config.predicate.test(bodyObj); - exchange.getAttributes().put(TEST_ATTRIBUTE, test); - return Flux.just(dataBuffer.factory().allocateBuffer().write(dataBuffer.asByteBuffer())); - }; - return messageReader.read(type, inputMessage, Collections.EMPTY_MAP).flatMap(mapper); - } - } - return Flux.just(dataBuffer.factory().allocateBuffer().write(dataBuffer.asByteBuffer())); - }); - BodyInserter bodyInserter = BodyInserters.fromPublisher(origBody, DataBuffer.class); - CachedBodyOutputMessage outputMessage = new CachedBodyOutputMessage(exchange, - exchange.getRequest().getHeaders()); - return bodyInserter.insert(outputMessage, new BodyInserterContext()) - // .log("modify_request", Level.INFO) - .then(Mono.defer(() -> { - boolean test = (Boolean) exchange.getAttributes() - .getOrDefault(TEST_ATTRIBUTE, Boolean.FALSE); - exchange.getAttributes().remove(TEST_ATTRIBUTE); - exchange.getAttributes().put(CACHED_REQUEST_BODY_KEY, - outputMessage.getBody()); - return Mono.just(test); - })); + return exchange.getRequest().getBody().collect(new DataBufferMapCollector()).flatMap(dataBufferMap -> { + BodyInserter bodyInserter = BodyInserters.fromPublisher(dataBufferMap.get(BODY_ONE), DataBuffer.class); + CachedBodyOutputMessage outputMessage = new CachedBodyOutputMessage(exchange, + exchange.getRequest().getHeaders()); + return bodyInserter.insert(outputMessage, new BodyInserterContext()) + // .log("modify_request", Level.INFO) + .then(Mono.defer(() -> { + ResolvableType type = ResolvableType.forClass(inClass); + for (HttpMessageReader messageReader : messageReaders) { + if (messageReader.canRead(type, exchange.getRequest().getHeaders().getContentType())) { + ReactiveHttpInputMessage inputMessage = new ReadBodyReactiveHttpInputMessage(dataBufferMap.get(BODY_TWO), + exchange.getRequest().getHeaders()); + Function mapper = (bodyObj) -> { + exchange.getAttributes().put(CACHE_REQUEST_BODY_OBJECT_KEY, bodyObj); + exchange.getAttributes().put(CACHED_REQUEST_BODY_KEY, + outputMessage.getBody()); + boolean test = config.predicate.test(bodyObj); + return Mono.just(test); + }; + return messageReader.readMono(type, inputMessage, Collections.EMPTY_MAP).flatMap(mapper); + } + } + return Mono.just(false); + })); + }); } }; } + /** + * This {@link Collector} is meant to collect the {@code Flux} from the request body into a {@link Map} + * which contains two copy of the body, one under the key {@code orig} and the other under the {@key copy}. + */ + class DataBufferMapCollector implements Collector>, Map>> { + public static final String BODY_ONE = "bodyOne"; + public static final String BODY_TWO = "bodyTwo"; + + @Override + public Supplier>> supplier() { + return () -> new HashMap>(); + } + + @Override + public BiConsumer>, DataBuffer> accumulator() { + return (dataBufferMap, dataBuffer) -> { + accumulate(BODY_ONE, dataBufferMap, dataBuffer); + accumulate(BODY_TWO, dataBufferMap, dataBuffer); + }; + } + + private void accumulate(String key, Map> dataBufferMap, DataBuffer dataBuffer) { + if (dataBufferMap.get(key) == null) { + dataBufferMap.put(key, Flux.just(copy(dataBuffer))); + } else { + dataBufferMap.put(key, dataBufferMap.get(key).mergeWith(Flux.just(copy(dataBuffer)))); + } + } + + @Override + public BinaryOperator>> combiner() { + return (map1, map2) -> { + map2.forEach((k, v) -> map1.merge(k, v, (v1, v2) -> v1.mergeWith(v2))); + return map1; + }; + } + + @Override + public Function>, Map>> finisher() { + return Function.identity(); + } + + @Override + public Set characteristics() { + return new HashSet(); + } + + private DataBuffer copy(DataBuffer dataBuffer) { + return dataBuffer.factory().allocateBuffer().write(dataBuffer.asByteBuffer()); + } + } + @Override @SuppressWarnings("unchecked") public Predicate apply(Config config) {