diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtils.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtils.java index 1091b6ed..2b6c85a9 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtils.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtils.java @@ -25,7 +25,7 @@ import java.util.Set; import java.util.function.Function; import java.util.function.Predicate; -import io.netty.buffer.EmptyByteBuf; +import io.netty.buffer.Unpooled; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; import reactor.core.publisher.Flux; @@ -36,9 +36,10 @@ import org.springframework.cloud.gateway.filter.factory.GatewayFilterFactory; import org.springframework.cloud.gateway.handler.AsyncPredicate; import org.springframework.cloud.gateway.handler.predicate.RoutePredicateFactory; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferFactory; import org.springframework.core.io.buffer.DataBufferUtils; +import org.springframework.core.io.buffer.DefaultDataBuffer; import org.springframework.core.io.buffer.NettyDataBuffer; -import org.springframework.core.io.buffer.NettyDataBufferFactory; import org.springframework.http.HttpStatus; import org.springframework.http.server.reactive.AbstractServerHttpResponse; import org.springframework.http.server.reactive.ServerHttpRequest; @@ -165,6 +166,8 @@ public final class ServerWebExchangeUtils { */ public static final String GATEWAY_LOADBALANCER_RESPONSE_ATTR = qualify("gatewayLoadBalancerResponse"); + private static final byte[] EMPTY_BYTES = {}; + private ServerWebExchangeUtils() { throw new AssertionError("Must not instantiate utility class."); } @@ -343,10 +346,9 @@ public final class ServerWebExchangeUtils { private static Mono cacheRequestBody(ServerWebExchange exchange, boolean cacheDecoratedRequest, Function> function) { ServerHttpResponse response = exchange.getResponse(); - NettyDataBufferFactory factory = (NettyDataBufferFactory) response.bufferFactory(); + DataBufferFactory factory = response.bufferFactory(); // Join all the DataBuffers so we have a single DataBuffer for the body - return DataBufferUtils.join(exchange.getRequest().getBody()) - .defaultIfEmpty(factory.wrap(new EmptyByteBuf(factory.getByteBufAllocator()))) + return DataBufferUtils.join(exchange.getRequest().getBody()).defaultIfEmpty(factory.wrap(EMPTY_BYTES)) .map(dataBuffer -> decorate(exchange, dataBuffer, cacheDecoratedRequest)) .switchIfEmpty(Mono.just(exchange.getRequest())).flatMap(function); } @@ -363,14 +365,23 @@ public final class ServerWebExchangeUtils { ServerHttpRequest decorator = new ServerHttpRequestDecorator(exchange.getRequest()) { @Override public Flux getBody() { - return Mono.fromSupplier(() -> { + return Mono.fromSupplier(() -> { if (exchange.getAttributeOrDefault(CACHED_REQUEST_BODY_ATTR, null) == null) { // probably == downstream closed or no body return null; } - // TODO: deal with Netty - NettyDataBuffer pdb = (NettyDataBuffer) dataBuffer; - return pdb.factory().wrap(pdb.getNativeBuffer().retainedSlice()); + if (dataBuffer instanceof NettyDataBuffer) { + NettyDataBuffer pdb = (NettyDataBuffer) dataBuffer; + return pdb.factory().wrap(pdb.getNativeBuffer().retainedSlice()); + } + else if (dataBuffer instanceof DefaultDataBuffer) { + DefaultDataBuffer ddf = (DefaultDataBuffer) dataBuffer; + return ddf.factory().wrap(Unpooled.wrappedBuffer(ddf.getNativeBuffer()).nioBuffer()); + } + else { + throw new IllegalArgumentException( + "Unable to handle DataBuffer of type " + dataBuffer.getClass()); + } }).flux(); } }; diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtilsTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtilsTests.java index f34b1f96..ed8bfcdd 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtilsTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/support/ServerWebExchangeUtilsTests.java @@ -23,10 +23,14 @@ import java.util.Map; import org.junit.Assert; import org.junit.Test; +import org.springframework.core.io.buffer.DefaultDataBuffer; import org.springframework.mock.http.server.reactive.MockServerHttpRequest; import org.springframework.mock.web.server.MockServerWebExchange; +import org.springframework.web.reactive.function.server.HandlerStrategies; +import org.springframework.web.reactive.function.server.ServerRequest; import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.CACHED_REQUEST_BODY_ATTR; import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.expand; public class ServerWebExchangeUtilsTests { @@ -52,6 +56,20 @@ public class ServerWebExchangeUtilsTests { Assert.assertThrows(IllegalArgumentException.class, () -> expand(exchange, "my-{foo}-{baz}")); } + @Test + public void defaultDataBufferHandling() { + MockServerWebExchange exchange = mockExchange(Collections.emptyMap()); + exchange.getAttributes().put(CACHED_REQUEST_BODY_ATTR, "foo"); + + ServerWebExchangeUtils + .cacheRequestBodyAndRequest(exchange, + (serverHttpRequest) -> ServerRequest + .create(exchange.mutate().request(serverHttpRequest).build(), + HandlerStrategies.withDefaults().messageReaders()) + .bodyToMono(DefaultDataBuffer.class)) + .block(); + } + private MockServerWebExchange mockExchange(Map vars) { MockServerHttpRequest request = MockServerHttpRequest.get("/get").build(); MockServerWebExchange exchange = MockServerWebExchange.from(request);