From e0a78c81f1a6a89488467102120535a2b6358ced Mon Sep 17 00:00:00 2001 From: fangfeikun <784945099@qq.com> Date: Sun, 19 Apr 2020 22:36:37 +0800 Subject: [PATCH] Releases body after sending modified body to downstream fails. Fixes gh-1520 --- .../rewrite/CachedBodyOutputMessage.java | 7 + ...ModifyRequestBodyGatewayFilterFactory.java | 15 +- ...dyGatewayFilterFactorySslTimeoutTests.java | 182 ++++++++++++++++++ 3 files changed, 203 insertions(+), 1 deletion(-) create mode 100644 spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactorySslTimeoutTests.java diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/CachedBodyOutputMessage.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/CachedBodyOutputMessage.java index 872a2764..cfd97597 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/CachedBodyOutputMessage.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/CachedBodyOutputMessage.java @@ -39,6 +39,8 @@ public class CachedBodyOutputMessage implements ReactiveHttpOutputMessage { private final HttpHeaders httpHeaders; + private boolean cached = false; + private Flux body = Flux.error(new IllegalStateException( "The body is not set. " + "Did handling complete with success?")); @@ -57,6 +59,10 @@ public class CachedBodyOutputMessage implements ReactiveHttpOutputMessage { return false; } + boolean isCached() { + return this.cached; + } + @Override public HttpHeaders getHeaders() { return this.httpHeaders; @@ -82,6 +88,7 @@ public class CachedBodyOutputMessage implements ReactiveHttpOutputMessage { public Mono writeWith(Publisher body) { this.body = Flux.from(body); + this.cached = true; return Mono.empty(); } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactory.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactory.java index 3c4b2dfc..653e9870 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactory.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactory.java @@ -18,6 +18,7 @@ package org.springframework.cloud.gateway.filter.factory.rewrite; import java.util.List; import java.util.Map; +import java.util.function.Function; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; @@ -27,6 +28,7 @@ import org.springframework.cloud.gateway.filter.GatewayFilterChain; import org.springframework.cloud.gateway.filter.factory.AbstractGatewayFilterFactory; import org.springframework.cloud.gateway.support.BodyInserterContext; import org.springframework.core.io.buffer.DataBuffer; +import org.springframework.core.io.buffer.DataBufferUtils; import org.springframework.http.HttpHeaders; import org.springframework.http.codec.HttpMessageReader; import org.springframework.http.codec.ServerCodecConfigurer; @@ -105,7 +107,9 @@ public class ModifyRequestBodyGatewayFilterFactory extends outputMessage); return chain .filter(exchange.mutate().request(decorator).build()); - })); + })).onErrorResume( + (Function>) throwable -> release( + exchange, outputMessage, throwable)); } @Override @@ -118,6 +122,15 @@ public class ModifyRequestBodyGatewayFilterFactory extends }; } + protected Mono release(ServerWebExchange exchange, + CachedBodyOutputMessage outputMessage, Throwable throwable) { + if (outputMessage.isCached()) { + return outputMessage.getBody().map(DataBufferUtils::release) + .then(Mono.error(throwable)); + } + return Mono.error(throwable); + } + ServerHttpRequestDecorator decorate(ServerWebExchange exchange, HttpHeaders headers, CachedBodyOutputMessage outputMessage) { return new ServerHttpRequestDecorator(exchange.getRequest()) { diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactorySslTimeoutTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactorySslTimeoutTests.java new file mode 100644 index 00000000..6e161982 --- /dev/null +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/factory/rewrite/ModifyRequestBodyGatewayFilterFactorySslTimeoutTests.java @@ -0,0 +1,182 @@ +/* + * Copyright 2013-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.cloud.gateway.filter.factory.rewrite; + +import java.util.concurrent.atomic.AtomicInteger; + +import javax.net.ssl.SSLException; + +import io.netty.handler.ssl.SslContext; +import io.netty.handler.ssl.SslContextBuilder; +import io.netty.handler.ssl.util.InsecureTrustManagerFactory; +import io.netty.util.internal.PlatformDependent; +import org.junit.Before; +import org.junit.Test; +import org.junit.runner.RunWith; +import reactor.core.publisher.Mono; +import reactor.netty.http.client.HttpClient; + +import org.springframework.beans.factory.annotation.Autowired; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.boot.SpringBootConfiguration; +import org.springframework.boot.autoconfigure.EnableAutoConfiguration; +import org.springframework.boot.test.context.SpringBootTest; +import org.springframework.cloud.gateway.route.RouteLocator; +import org.springframework.cloud.gateway.route.builder.RouteLocatorBuilder; +import org.springframework.cloud.gateway.test.BaseWebClientTests; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.DependsOn; +import org.springframework.context.annotation.Import; +import org.springframework.context.annotation.Primary; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.MediaType; +import org.springframework.http.client.reactive.ReactorClientHttpConnector; +import org.springframework.http.codec.ServerCodecConfigurer; +import org.springframework.test.annotation.DirtiesContext; +import org.springframework.test.context.ActiveProfiles; +import org.springframework.test.context.junit4.SpringRunner; +import org.springframework.web.reactive.function.BodyInserters; +import org.springframework.web.server.ServerWebExchange; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.boot.test.context.SpringBootTest.WebEnvironment.RANDOM_PORT; + +/** + * @author fangfeikun + */ +@RunWith(SpringRunner.class) +@SpringBootTest(webEnvironment = RANDOM_PORT, + properties = { "spring.cloud.gateway.httpclient.ssl.handshake-timeout=1ms", + "spring.main.allow-bean-definition-overriding=true" }) +@DirtiesContext +@ActiveProfiles("single-cert-ssl") +public class ModifyRequestBodyGatewayFilterFactorySslTimeoutTests + extends BaseWebClientTests { + + @Autowired + AtomicInteger releaseCount; + + @Before + public void setup() { + try { + SslContext sslContext = SslContextBuilder.forClient() + .trustManager(InsecureTrustManagerFactory.INSTANCE).build(); + HttpClient httpClient = HttpClient.create() + .secure(ssl -> ssl.sslContext(sslContext)); + setup(new ReactorClientHttpConnector(httpClient), + "https://localhost:" + port); + } + catch (SSLException e) { + throw new RuntimeException(e); + } + } + + @Test + public void modifyRequestBodySSLTimeout() { + testClient.post().uri("/post") + .header("Host", "www.modifyrequestbodyssltimeout.org") + .header(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_XML_VALUE) + .body(BodyInserters.fromValue("request")).exchange().expectStatus() + .isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR).expectBody() + .jsonPath("message").isEqualTo("handshake timed out after 1ms"); + } + + @Test + public void modifyRequestBodyRelease() { + releaseCount.set(0); + // long initialUsedDirectMemory = PlatformDependent.usedDirectMemory(); + for (int i = 0; i < 10; i++) { + testClient.post().uri("/post") + .header("Host", "www.modifyrequestbodyssltimeout.org") + .header(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_XML_VALUE) + .body(BodyInserters.fromValue("request")).exchange().expectStatus() + .isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR); + long usedDirectMemory = PlatformDependent.usedDirectMemory(); + // Assert.assertTrue(usedDirectMemory - initialUsedDirectMemory < 2 * 10 * 10 + // * 1024 * 1024); + } + assertThat(releaseCount).hasValue(10); + } + + @Test + public void modifyRequestBodyHappenedError() { + testClient.post().uri("/post") + .header("Host", "www.modifyrequestbodyexception.org") + .header(HttpHeaders.CONTENT_TYPE, MediaType.APPLICATION_XML_VALUE) + .body(BodyInserters.fromValue("request")).exchange().expectStatus() + .isEqualTo(HttpStatus.INTERNAL_SERVER_ERROR).expectBody() + .jsonPath("message").isEqualTo("modify body exception"); + } + + @EnableAutoConfiguration + @SpringBootConfiguration(proxyBeanMethods = false) + @Import(DefaultTestConfig.class) + public static class TestConfig { + + @Value("${test.uri}") + String uri; + + @Bean + @DependsOn("testModifyRequestBodyGatewayFilterFactory") + public RouteLocator testRouteLocator(RouteLocatorBuilder builder) { + return builder.routes().route("test_modify_request_body_ssl_timeout", + r -> r.order(-1).host("**.modifyrequestbodyssltimeout.org") + .filters(f -> f.modifyRequestBody(String.class, String.class, + MediaType.APPLICATION_JSON_VALUE, + (serverWebExchange, aVoid) -> { + byte[] largeBody = new byte[10 * 1024 * 1024]; + return Mono.just(new String(largeBody)); + })) + .uri(uri)) + .route("test_modify_request_body_exception", r -> r.order(-1) + .host("**.modifyrequestbodyexception.org") + .filters(f -> f.modifyRequestBody(String.class, String.class, + MediaType.APPLICATION_JSON_VALUE, + (serverWebExchange, body) -> { + return Mono.error( + new Exception("modify body exception")); + })) + .uri(uri)) + .build(); + } + + @Bean + public AtomicInteger count() { + return new AtomicInteger(); + } + + @Bean + @Primary + public ModifyRequestBodyGatewayFilterFactory testModifyRequestBodyGatewayFilterFactory( + ServerCodecConfigurer codecConfigurer, AtomicInteger count) { + return new ModifyRequestBodyGatewayFilterFactory( + codecConfigurer.getReaders()) { + @Override + protected Mono release(ServerWebExchange exchange, + CachedBodyOutputMessage outputMessage, Throwable throwable) { + if (outputMessage.isCached()) { + count.incrementAndGet(); + } + return super.release(exchange, outputMessage, throwable); + } + }; + } + + } + +}