From dce413ae66018b0b802dff170fa2eca9d6abc3dc Mon Sep 17 00:00:00 2001 From: Olga Maciaszek-Sharma Date: Wed, 9 Dec 2020 06:37:58 -0600 Subject: [PATCH] Do not pass response body to LB lifecycle beans. (#864) * Do not pass response body to LB lifecycle beans. * Fix argument name. --- .../client/loadbalancer/DefaultResponse.java | 19 +++ .../loadbalancer/HttpRequestContext.java | 38 ------ .../client/loadbalancer/RequestData.java | 114 ++++++++++++++++++ ...stContext.java => RequestDataContext.java} | 19 +-- .../client/loadbalancer/ResponseData.java | 100 +++++++++++++++ .../RetryLoadBalancerInterceptor.java | 4 +- .../ServerHttpRequestContext.java | 44 ------- ...torLoadBalancerExchangeFilterFunction.java | 16 +-- ...bleLoadBalancerExchangeFilterFunction.java | 11 +- ...adBalancerExchangeFilterFunctionTests.java | 17 +-- ...xchangeFilterFunctionIntegrationTests.java | 8 +- ...ckySessionServiceInstanceListSupplier.java | 35 ++---- ...ssionServiceInstanceListSupplierTests.java | 57 ++------- 13 files changed, 293 insertions(+), 189 deletions(-) delete mode 100644 spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/HttpRequestContext.java create mode 100644 spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestData.java rename spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/{ClientRequestContext.java => RequestDataContext.java} (62%) create mode 100644 spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ResponseData.java delete mode 100644 spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ServerHttpRequestContext.java diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/DefaultResponse.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/DefaultResponse.java index 5b1e6ea8..85c74177 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/DefaultResponse.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/DefaultResponse.java @@ -16,6 +16,8 @@ package org.springframework.cloud.client.loadbalancer; +import java.util.Objects; + import org.springframework.cloud.client.ServiceInstance; import org.springframework.core.style.ToStringCreator; @@ -53,4 +55,21 @@ public class DefaultResponse implements Response { return to.toString(); } + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (!(o instanceof DefaultResponse)) { + return false; + } + DefaultResponse that = (DefaultResponse) o; + return Objects.equals(serviceInstance, that.serviceInstance); + } + + @Override + public int hashCode() { + return Objects.hash(serviceInstance); + } + } diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/HttpRequestContext.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/HttpRequestContext.java deleted file mode 100644 index 7f9c1d12..00000000 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/HttpRequestContext.java +++ /dev/null @@ -1,38 +0,0 @@ -/* - * Copyright 2012-2020 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.client.loadbalancer; - -import org.springframework.http.HttpRequest; - -/** - * @author Olga Maciaszek-Sharma - */ -public class HttpRequestContext extends DefaultRequestContext { - - public HttpRequestContext(HttpRequest httpRequest) { - this(httpRequest, "default"); - } - - public HttpRequestContext(HttpRequest httpRequest, String hint) { - super(httpRequest, hint); - } - - public HttpRequest getClientRequest() { - return (HttpRequest) super.getClientRequest(); - } - -} diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestData.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestData.java new file mode 100644 index 00000000..bb7fc33e --- /dev/null +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestData.java @@ -0,0 +1,114 @@ +/* + * Copyright 2012-2020 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.client.loadbalancer; + +import java.net.URI; +import java.util.HashMap; +import java.util.Map; +import java.util.Objects; + +import org.springframework.core.style.ToStringCreator; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpMethod; +import org.springframework.http.HttpRequest; +import org.springframework.util.MultiValueMap; +import org.springframework.web.reactive.function.client.ClientRequest; + +/** + * Represents the data of the request that can be safely read (without passing request reactive stream values). + * + * @author Olga Maciaszek-Sharma + * @since 3.0.0 + */ +public class RequestData { + + private final HttpMethod httpMethod; + + private final URI url; + + private final HttpHeaders headers; + + private final MultiValueMap cookies; + + private final Map attributes; + + public RequestData(HttpMethod httpMethod, URI url, HttpHeaders headers, MultiValueMap cookies, + Map attributes) { + this.httpMethod = httpMethod; + this.url = url; + this.headers = headers; + this.cookies = cookies; + this.attributes = attributes; + } + + public RequestData(ClientRequest request) { + this(request.method(), request.url(), request.headers(), request.cookies(), request.attributes()); + } + + public RequestData(HttpRequest request) { + this(request.getMethod(), request.getURI(), request.getHeaders(), null, new HashMap<>()); + } + + public HttpMethod getHttpMethod() { + return httpMethod; + } + + public URI getUrl() { + return url; + } + + public HttpHeaders getHeaders() { + return headers; + } + + public MultiValueMap getCookies() { + return cookies; + } + + public Map getAttributes() { + return attributes; + } + + @Override + public String toString() { + ToStringCreator to = new ToStringCreator(this); + to.append("httpMethod", httpMethod); + to.append("url", url); + to.append("headers", headers); + to.append("cookies", cookies); + return to.toString(); + } + + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (!(o instanceof RequestData)) { + return false; + } + RequestData that = (RequestData) o; + return httpMethod == that.httpMethod && Objects.equals(url, that.url) && Objects.equals(headers, that.headers) + && Objects.equals(cookies, that.cookies) && Objects.equals(attributes, that.attributes); + } + + @Override + public int hashCode() { + return Objects.hash(httpMethod, url, headers, cookies, attributes); + } + +} diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ClientRequestContext.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestDataContext.java similarity index 62% rename from spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ClientRequestContext.java rename to spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestDataContext.java index 8bdbab38..003aa9f6 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ClientRequestContext.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RequestDataContext.java @@ -17,28 +17,29 @@ package org.springframework.cloud.client.loadbalancer; import org.springframework.http.HttpMethod; -import org.springframework.web.reactive.function.client.ClientRequest; /** + * A {@link RequestData}-based {@link DefaultRequestContext}. + * * @author Olga Maciaszek-Sharma * @since 3.0.0 */ -public class ClientRequestContext extends DefaultRequestContext { +public class RequestDataContext extends DefaultRequestContext { - public ClientRequestContext(ClientRequest clientRequest) { - this(clientRequest, "default"); + public RequestDataContext(RequestData requestData) { + this(requestData, "default"); } - public ClientRequestContext(ClientRequest clientRequest, String hint) { - super(clientRequest, hint); + public RequestDataContext(RequestData requestData, String hint) { + super(requestData, hint); } - public ClientRequest getClientRequest() { - return (ClientRequest) super.getClientRequest(); + public RequestData getClientRequest() { + return (RequestData) super.getClientRequest(); } public HttpMethod method() { - return ((ClientRequest) super.getClientRequest()).method(); + return ((RequestData) super.getClientRequest()).getHttpMethod(); } } diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ResponseData.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ResponseData.java new file mode 100644 index 00000000..8405c35e --- /dev/null +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ResponseData.java @@ -0,0 +1,100 @@ +/* + * Copyright 2012-2020 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.client.loadbalancer; + +import java.util.Objects; + +import org.springframework.core.style.ToStringCreator; +import org.springframework.http.HttpHeaders; +import org.springframework.http.HttpStatus; +import org.springframework.http.ResponseCookie; +import org.springframework.util.MultiValueMap; +import org.springframework.web.reactive.function.client.ClientResponse; + +/** + * Represents the data of the request that can be safely read (without passing request reactive stream values). + * + * @author Olga Maciaszek-Sharma + * @since 3.0.0 + */ +public class ResponseData { + + private final HttpStatus httpStatus; + + private final HttpHeaders headers; + + private final MultiValueMap cookies; + + private final RequestData requestData; + + public ResponseData(HttpStatus httpStatus, HttpHeaders headers, MultiValueMap cookies, + RequestData requestData) { + this.httpStatus = httpStatus; + this.headers = headers; + this.cookies = cookies; + this.requestData = requestData; + } + + public ResponseData(ClientResponse response, RequestData requestData) { + httpStatus = response.statusCode(); + headers = response.headers().asHttpHeaders(); + cookies = response.cookies(); + this.requestData = requestData; + } + + public HttpStatus getHttpStatus() { + return httpStatus; + } + + public HttpHeaders getHeaders() { + return headers; + } + + public MultiValueMap getCookies() { + return cookies; + } + + public RequestData getRequestData() { + return requestData; + } + + @Override + public String toString() { + ToStringCreator to = new ToStringCreator(this); + to.append("httpStatus", httpStatus); + return to.toString(); + } + + @Override + public boolean equals(Object o) { + if (this == o) { + return true; + } + if (!(o instanceof ResponseData)) { + return false; + } + ResponseData that = (ResponseData) o; + return httpStatus == that.httpStatus && Objects.equals(headers, that.headers) + && Objects.equals(cookies, that.cookies) && Objects.equals(requestData, that.requestData); + } + + @Override + public int hashCode() { + return Objects.hash(httpStatus, headers, cookies, requestData); + } + +} diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RetryLoadBalancerInterceptor.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RetryLoadBalancerInterceptor.java index 46b44a5f..48038d3f 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RetryLoadBalancerInterceptor.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/RetryLoadBalancerInterceptor.java @@ -90,7 +90,7 @@ public class RetryLoadBalancerInterceptor implements ClientHttpRequestIntercepto Set supportedLifecycleProcessors = LoadBalancerLifecycleValidator .getSupportedLifecycleProcessors( loadBalancerFactory.getInstances(serviceName, LoadBalancerLifecycle.class), - HttpRequestContext.class, ClientHttpResponse.class, ServiceInstance.class); + RequestDataContext.class, ResponseData.class, ServiceInstance.class); if (serviceInstance == null) { if (LOG.isDebugEnabled()) { LOG.debug("Service instance retrieved from LoadBalancedRetryContext: was null. " @@ -103,7 +103,7 @@ public class RetryLoadBalancerInterceptor implements ClientHttpRequestIntercepto } String hint = getHint(serviceName); DefaultRequest lbRequest = new DefaultRequest<>( - new RetryableRequestContext(previousServiceInstance, request, hint)); + new RetryableRequestContext(previousServiceInstance, new RequestData(request), hint)); supportedLifecycleProcessors.forEach(lifecycle -> lifecycle.onStart(lbRequest)); serviceInstance = loadBalancer.choose(serviceName, lbRequest); if (LOG.isDebugEnabled()) { diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ServerHttpRequestContext.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ServerHttpRequestContext.java deleted file mode 100644 index 99f0d2be..00000000 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/ServerHttpRequestContext.java +++ /dev/null @@ -1,44 +0,0 @@ -/* - * Copyright 2012-2020 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.client.loadbalancer; - -import org.springframework.http.HttpMethod; -import org.springframework.http.server.reactive.ServerHttpRequest; - -/** - * @author Olga Maciaszek-Sharma - * @since 3.0.0 - */ -public class ServerHttpRequestContext extends DefaultRequestContext { - - public ServerHttpRequestContext(ServerHttpRequest serverHttpRequest) { - this(serverHttpRequest, "default"); - } - - public ServerHttpRequestContext(ServerHttpRequest serverHttpRequest, String hint) { - super(serverHttpRequest, hint); - } - - public ServerHttpRequest getClientRequest() { - return (ServerHttpRequest) super.getClientRequest(); - } - - public HttpMethod method() { - return ((ServerHttpRequest) super.getClientRequest()).getMethod(); - } - -} diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunction.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunction.java index 7ce1df5f..a12ce8d2 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunction.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunction.java @@ -24,14 +24,16 @@ import org.apache.commons.logging.LogFactory; import reactor.core.publisher.Mono; import org.springframework.cloud.client.ServiceInstance; -import org.springframework.cloud.client.loadbalancer.ClientRequestContext; import org.springframework.cloud.client.loadbalancer.CompletionContext; import org.springframework.cloud.client.loadbalancer.DefaultRequest; import org.springframework.cloud.client.loadbalancer.EmptyResponse; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycle; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycleValidator; import org.springframework.cloud.client.loadbalancer.Request; +import org.springframework.cloud.client.loadbalancer.RequestData; +import org.springframework.cloud.client.loadbalancer.RequestDataContext; import org.springframework.cloud.client.loadbalancer.Response; +import org.springframework.cloud.client.loadbalancer.ResponseData; import org.springframework.http.HttpStatus; import org.springframework.web.reactive.function.client.ClientRequest; import org.springframework.web.reactive.function.client.ClientResponse; @@ -78,10 +80,10 @@ public class ReactorLoadBalancerExchangeFilterFunction implements LoadBalancedEx Set supportedLifecycleProcessors = LoadBalancerLifecycleValidator .getSupportedLifecycleProcessors( loadBalancerFactory.getInstances(serviceId, LoadBalancerLifecycle.class), - ClientRequestContext.class, ClientResponse.class, ServiceInstance.class); + RequestDataContext.class, ResponseData.class, ServiceInstance.class); String hint = getHint(serviceId, properties.getHint()); - DefaultRequest lbRequest = new DefaultRequest<>( - new ClientRequestContext(clientRequest, hint)); + RequestData requestData = new RequestData(clientRequest); + DefaultRequest lbRequest = new DefaultRequest<>(new RequestDataContext(requestData, hint)); supportedLifecycleProcessors.forEach(lifecycle -> lifecycle.onStart(lbRequest)); return choose(serviceId, lbRequest).flatMap(lbResponse -> { ServiceInstance instance = lbResponse.getServer(); @@ -106,15 +108,15 @@ public class ReactorLoadBalancerExchangeFilterFunction implements LoadBalancedEx stickySessionProperties.isAddServiceInstanceCookie()); return next.exchange(newRequest) .doOnError(throwable -> supportedLifecycleProcessors.forEach( - lifecycle -> lifecycle.onComplete(new CompletionContext( + lifecycle -> lifecycle.onComplete(new CompletionContext( CompletionContext.Status.FAILED, throwable, lbResponse)))) .doOnSuccess(clientResponse -> supportedLifecycleProcessors.forEach( lifecycle -> lifecycle.onComplete(new CompletionContext<>(CompletionContext.Status.SUCCESS, - lbResponse, clientResponse)))); + lbResponse, new ResponseData(clientResponse, requestData))))); }); } - protected Mono> choose(String serviceId, Request request) { + protected Mono> choose(String serviceId, Request request) { ReactiveLoadBalancer loadBalancer = loadBalancerFactory.getInstance(serviceId); if (loadBalancer == null) { return Mono.just(new EmptyResponse()); diff --git a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunction.java b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunction.java index abb7c23e..880eee7b 100644 --- a/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunction.java +++ b/spring-cloud-commons/src/main/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunction.java @@ -31,14 +31,16 @@ import reactor.util.retry.Retry; import reactor.util.retry.RetrySpec; import org.springframework.cloud.client.ServiceInstance; -import org.springframework.cloud.client.loadbalancer.ClientRequestContext; import org.springframework.cloud.client.loadbalancer.CompletionContext; import org.springframework.cloud.client.loadbalancer.DefaultRequest; import org.springframework.cloud.client.loadbalancer.EmptyResponse; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycle; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycleValidator; import org.springframework.cloud.client.loadbalancer.Request; +import org.springframework.cloud.client.loadbalancer.RequestData; +import org.springframework.cloud.client.loadbalancer.RequestDataContext; import org.springframework.cloud.client.loadbalancer.Response; +import org.springframework.cloud.client.loadbalancer.ResponseData; import org.springframework.cloud.client.loadbalancer.RetryableRequestContext; import org.springframework.http.HttpStatus; import org.springframework.web.reactive.function.client.ClientRequest; @@ -99,10 +101,11 @@ public class RetryableLoadBalancerExchangeFilterFunction implements LoadBalanced Set supportedLifecycleProcessors = LoadBalancerLifecycleValidator .getSupportedLifecycleProcessors( loadBalancerFactory.getInstances(serviceId, LoadBalancerLifecycle.class), - ClientRequestContext.class, ClientResponse.class, ServiceInstance.class); + RequestDataContext.class, ResponseData.class, ServiceInstance.class); String hint = getHint(serviceId, properties.getHint()); + RequestData requestData = new RequestData(clientRequest); DefaultRequest lbRequest = new DefaultRequest<>( - new RetryableRequestContext(null, clientRequest, hint)); + new RetryableRequestContext(null, requestData, hint)); supportedLifecycleProcessors.forEach(lifecycle -> lifecycle.onStart(lbRequest)); return Mono.defer(() -> choose(serviceId, lbRequest).flatMap(lbResponse -> { ServiceInstance instance = lbResponse.getServer(); @@ -132,7 +135,7 @@ public class RetryableLoadBalancerExchangeFilterFunction implements LoadBalanced CompletionContext.Status.FAILED, throwable, lbResponse)))) .doOnSuccess(clientResponse -> supportedLifecycleProcessors.forEach( lifecycle -> lifecycle.onComplete(new CompletionContext<>(CompletionContext.Status.SUCCESS, - lbResponse, clientResponse)))) + lbResponse, new ResponseData(clientResponse, requestData))))) .map(clientResponse -> { loadBalancerRetryContext.setClientResponse(clientResponse); if (shouldRetrySameServiceInstance(loadBalancerRetryContext)) { diff --git a/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunctionTests.java b/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunctionTests.java index 30887182..45741e15 100644 --- a/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunctionTests.java +++ b/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/ReactorLoadBalancerExchangeFilterFunctionTests.java @@ -42,6 +42,7 @@ import org.springframework.cloud.client.loadbalancer.CompletionContext; import org.springframework.cloud.client.loadbalancer.DefaultRequestContext; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycle; import org.springframework.cloud.client.loadbalancer.Request; +import org.springframework.cloud.client.loadbalancer.ResponseData; import org.springframework.context.annotation.Bean; import org.springframework.http.HttpStatus; import org.springframework.web.bind.annotation.RequestMapping; @@ -94,8 +95,8 @@ class ReactorLoadBalancerExchangeFilterFunctionTests { @Test void correctResponseReturnedForExistingHostAndInstancePresent() { - ClientResponse clientResponse = WebClient.builder().baseUrl("http://testservice") - .filter(this.loadBalancerFunction).build().get().uri("/hello").exchange().block(); + ClientResponse clientResponse = WebClient.builder().baseUrl("http://testservice").filter(loadBalancerFunction) + .build().get().uri("/hello").exchange().block(); then(clientResponse.statusCode()).isEqualTo(HttpStatus.OK); then(clientResponse.bodyToMono(String.class).block()).isEqualTo("Hello World"); } @@ -118,7 +119,7 @@ class ReactorLoadBalancerExchangeFilterFunctionTests { @Test void exceptionNotThrownWhenFactoryReturnsNullLifecycleProcessorsMap() { assertThatCode(() -> WebClient.builder().baseUrl("http://serviceWithNoLifecycleProcessors") - .filter(this.loadBalancerFunction).build().get().uri("/hello").exchange().block()) + .filter(loadBalancerFunction).build().get().uri("/hello").exchange().block()) .doesNotThrowAnyException(); } @@ -127,8 +128,8 @@ class ReactorLoadBalancerExchangeFilterFunctionTests { final String callbackTestHint = "callbackTestHint"; loadBalancerProperties.getHint().put("testservice", "callbackTestHint"); final String result = "callbackTestResult"; - ClientResponse clientResponse = WebClient.builder().baseUrl("http://testservice") - .filter(this.loadBalancerFunction).build().get().uri("/callback").exchange().block(); + ClientResponse clientResponse = WebClient.builder().baseUrl("http://testservice").filter(loadBalancerFunction) + .build().get().uri("/callback").exchange().block(); Collection> lifecycleLogRequests = ((TestLoadBalancerLifecycle) factory .getInstances("testservice", LoadBalancerLifecycle.class).get("loadBalancerLifecycle")).getStartLog() @@ -140,9 +141,9 @@ class ReactorLoadBalancerExchangeFilterFunctionTests { assertThat(lifecycleLogRequests).extracting(request -> ((DefaultRequestContext) request.getContext()).getHint()) .contains(callbackTestHint); assertThat(anotherLifecycleLogRequests) - .extracting(completionContext -> ((ClientResponse) completionContext.getClientResponse()) - .bodyToMono(String.class).block()) - .contains(result); + .extracting(completionContext -> ((ResponseData) completionContext.getClientResponse()).getRequestData() + .getUrl().toString()) + .contains("http://testservice/callback"); } @SuppressWarnings({ "unchecked", "rawtypes" }) diff --git a/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunctionIntegrationTests.java b/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunctionIntegrationTests.java index f03d6e56..1b930324 100644 --- a/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunctionIntegrationTests.java +++ b/spring-cloud-commons/src/test/java/org/springframework/cloud/client/loadbalancer/reactive/RetryableLoadBalancerExchangeFilterFunctionIntegrationTests.java @@ -44,7 +44,9 @@ import org.springframework.cloud.client.loadbalancer.CompletionContext; import org.springframework.cloud.client.loadbalancer.DefaultRequestContext; import org.springframework.cloud.client.loadbalancer.LoadBalancerLifecycle; import org.springframework.cloud.client.loadbalancer.Request; +import org.springframework.cloud.client.loadbalancer.ResponseData; import org.springframework.context.annotation.Bean; +import org.springframework.http.HttpMethod; import org.springframework.http.HttpStatus; import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.RestController; @@ -112,9 +114,9 @@ class RetryableLoadBalancerExchangeFilterFunctionIntegrationTests { assertThat(lifecycleLogRequests).extracting(request -> ((DefaultRequestContext) request.getContext()).getHint()) .contains(callbackTestHint); assertThat(anotherLifecycleLogRequests) - .extracting(completionContext -> ((ClientResponse) completionContext.getClientResponse()) - .bodyToMono(String.class).block()) - .contains(result); + .extracting(completionContext -> ((ResponseData) completionContext.getClientResponse()).getRequestData() + .getHttpMethod()) + .contains(HttpMethod.GET); } @Test diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplier.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplier.java index 4022b00a..9f02f496 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplier.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplier.java @@ -24,13 +24,10 @@ import org.apache.commons.logging.LogFactory; import reactor.core.publisher.Flux; import org.springframework.cloud.client.ServiceInstance; -import org.springframework.cloud.client.loadbalancer.ClientRequestContext; import org.springframework.cloud.client.loadbalancer.Request; -import org.springframework.cloud.client.loadbalancer.ServerHttpRequestContext; +import org.springframework.cloud.client.loadbalancer.RequestDataContext; import org.springframework.cloud.client.loadbalancer.reactive.LoadBalancerProperties; -import org.springframework.http.HttpCookie; -import org.springframework.http.server.reactive.ServerHttpRequest; -import org.springframework.web.reactive.function.client.ClientRequest; +import org.springframework.util.MultiValueMap; /** * A session cookie based implementation of {@link ServiceInstanceListSupplier} that gives @@ -66,10 +63,13 @@ public class RequestBasedStickySessionServiceInstanceListSupplier extends Delega public Flux> get(Request request) { String instanceIdCookieName = properties.getStickySession().getInstanceIdCookieName(); Object context = request.getContext(); - if ((context instanceof ClientRequestContext)) { - ClientRequest originalRequest = ((ClientRequestContext) context).getClientRequest(); + if ((context instanceof RequestDataContext)) { + MultiValueMap cookies = ((RequestDataContext) context).getClientRequest().getCookies(); + if (cookies == null) { + return get(); + } // We expect there to be one value in this cookie - String cookie = originalRequest.cookies().getFirst(instanceIdCookieName); + String cookie = cookies.getFirst(instanceIdCookieName); if (cookie != null) { return get().map(serviceInstances -> selectInstance(serviceInstances, cookie)); } @@ -78,23 +78,8 @@ public class RequestBasedStickySessionServiceInstanceListSupplier extends Delega } return get(); } - if ((context instanceof ServerHttpRequestContext)) { - ServerHttpRequest originalRequest = ((ServerHttpRequestContext) context).getClientRequest(); - HttpCookie cookie = originalRequest.getCookies().getFirst(instanceIdCookieName); - if (cookie != null) { - return get().map(serviceInstances -> selectInstance(serviceInstances, cookie.getValue())); - } - if (LOG.isDebugEnabled()) { - LOG.debug("Cookie not found. Returning all instances returned by delegate."); - } - return get(); - } - if (LOG.isDebugEnabled()) { - LOG.debug("Searching for instances based on cookie not supported for ClientRequestContext type." - + " Returning all instances returned by delegate."); - } - // If no cookie is available, we return all the instances provided by the - // delegate. + // If the object type is not RequestData, we return all the instances provided by + // the delegate. return get(); } diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplierTests.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplierTests.java index 3a4df663..358c476d 100644 --- a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplierTests.java +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/RequestBasedStickySessionServiceInstanceListSupplierTests.java @@ -25,17 +25,14 @@ import reactor.core.publisher.Flux; import org.springframework.cloud.client.DefaultServiceInstance; import org.springframework.cloud.client.ServiceInstance; -import org.springframework.cloud.client.loadbalancer.ClientRequestContext; import org.springframework.cloud.client.loadbalancer.DefaultRequest; import org.springframework.cloud.client.loadbalancer.DefaultRequestContext; import org.springframework.cloud.client.loadbalancer.Request; -import org.springframework.cloud.client.loadbalancer.ServerHttpRequestContext; +import org.springframework.cloud.client.loadbalancer.RequestData; +import org.springframework.cloud.client.loadbalancer.RequestDataContext; import org.springframework.cloud.client.loadbalancer.reactive.LoadBalancerProperties; -import org.springframework.http.HttpCookie; import org.springframework.http.HttpHeaders; import org.springframework.http.server.reactive.ServerHttpRequest; -import org.springframework.util.LinkedMultiValueMap; -import org.springframework.util.MultiValueMap; import org.springframework.web.reactive.function.client.ClientRequest; import static org.assertj.core.api.Assertions.assertThat; @@ -77,7 +74,8 @@ class RequestBasedStickySessionServiceInstanceListSupplierTests { HttpHeaders headers = new HttpHeaders(); headers.add(properties.getStickySession().getInstanceIdCookieName(), "test-1"); when(clientRequest.cookies()).thenReturn(headers); - Request request = new DefaultRequest<>(new ClientRequestContext(clientRequest)); + Request request = new DefaultRequest<>( + new RequestDataContext(new RequestData(clientRequest))); List serviceInstances = supplier.get(request).blockFirst(); @@ -90,7 +88,8 @@ class RequestBasedStickySessionServiceInstanceListSupplierTests { HttpHeaders headers = new HttpHeaders(); headers.add(properties.getStickySession().getInstanceIdCookieName(), "test-4"); when(clientRequest.cookies()).thenReturn(headers); - Request request = new DefaultRequest<>(new ClientRequestContext(clientRequest)); + Request request = new DefaultRequest<>( + new RequestDataContext(new RequestData(clientRequest))); List serviceInstances = supplier.get(request).blockFirst(); @@ -100,48 +99,8 @@ class RequestBasedStickySessionServiceInstanceListSupplierTests { @Test void shouldReturnAllInstancesFromDelegateIfClientRequestHasNoCookie() { when(clientRequest.cookies()).thenReturn(new HttpHeaders()); - Request request = new DefaultRequest<>(new ClientRequestContext(clientRequest)); - - List serviceInstances = supplier.get(request).blockFirst(); - - assertThat(serviceInstances).hasSize(3); - } - - @Test - void shouldReturnInstanceBasedOnCookieFromServerHttpRequest() { - MultiValueMap cookies = new LinkedMultiValueMap<>(); - cookies.add(properties.getStickySession().getInstanceIdCookieName(), - new HttpCookie(properties.getStickySession().getInstanceIdCookieName(), "test-1")); - when(serverHttpRequest.getCookies()).thenReturn(cookies); - Request request = new DefaultRequest<>( - new ServerHttpRequestContext(serverHttpRequest)); - - List serviceInstances = supplier.get(request).blockFirst(); - - assertThat(serviceInstances).hasSize(1); - assertThat(serviceInstances.get(0).getInstanceId()).isEqualTo("test-1"); - } - - @Test - void shouldReturnAllDelegateInstancesIfInstanceBasedOnCookieFromServerHttpRequestNotFound() { - MultiValueMap cookies = new LinkedMultiValueMap<>(); - cookies.add(properties.getStickySession().getInstanceIdCookieName(), - new HttpCookie(properties.getStickySession().getInstanceIdCookieName(), "test-4")); - when(serverHttpRequest.getCookies()).thenReturn(cookies); - Request request = new DefaultRequest<>( - new ServerHttpRequestContext(serverHttpRequest)); - - List serviceInstances = supplier.get(request).blockFirst(); - - assertThat(serviceInstances).hasSize(3); - } - - @Test - void shouldReturnAllInstancesFromDelegateIfServerHttpRequestHasNoCookie() { - MultiValueMap cookies = new LinkedMultiValueMap<>(); - when(serverHttpRequest.getCookies()).thenReturn(cookies); - Request request = new DefaultRequest<>( - new ServerHttpRequestContext(serverHttpRequest)); + Request request = new DefaultRequest<>( + new RequestDataContext(new RequestData(clientRequest))); List serviceInstances = supplier.get(request).blockFirst();