Use a Collector to collect the request body into a map containing two copies.

This commit is contained in:
Ryan Baxter
2018-08-17 20:14:19 -04:00
parent a0508bbe1e
commit 31ac60d197

View File

@@ -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<DataBuffer> 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<DataBuffer>} 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<DataBuffer, Map<String, Flux<DataBuffer>>, Map<String, Flux<DataBuffer>>> {
public static final String BODY_ONE = "bodyOne";
public static final String BODY_TWO = "bodyTwo";
@Override
public Supplier<Map<String, Flux<DataBuffer>>> supplier() {
return () -> new HashMap<String, Flux<DataBuffer>>();
}
@Override
public BiConsumer<Map<String, Flux<DataBuffer>>, DataBuffer> accumulator() {
return (dataBufferMap, dataBuffer) -> {
accumulate(BODY_ONE, dataBufferMap, dataBuffer);
accumulate(BODY_TWO, dataBufferMap, dataBuffer);
};
}
private void accumulate(String key, Map<String, Flux<DataBuffer>> 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<Map<String, Flux<DataBuffer>>> combiner() {
return (map1, map2) -> {
map2.forEach((k, v) -> map1.merge(k, v, (v1, v2) -> v1.mergeWith(v2)));
return map1;
};
}
@Override
public Function<Map<String, Flux<DataBuffer>>, Map<String, Flux<DataBuffer>>> finisher() {
return Function.identity();
}
@Override
public Set<Characteristics> characteristics() {
return new HashSet<Characteristics>();
}
private DataBuffer copy(DataBuffer dataBuffer) {
return dataBuffer.factory().allocateBuffer().write(dataBuffer.asByteBuffer());
}
}
@Override
@SuppressWarnings("unchecked")
public Predicate<ServerWebExchange> apply(Config config) {