Do not pass response body to LB lifecycle beans. (#864)

* Do not pass response body to LB lifecycle beans.

* Fix argument name.
This commit is contained in:
Olga Maciaszek-Sharma
2020-12-09 06:37:58 -06:00
committed by GitHub
parent 0f98419c47
commit dce413ae66
13 changed files with 293 additions and 189 deletions

View File

@@ -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<ServiceInstance> {
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);
}
}

View File

@@ -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();
}
}

View File

@@ -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<String, String> cookies;
private final Map<String, Object> attributes;
public RequestData(HttpMethod httpMethod, URI url, HttpHeaders headers, MultiValueMap<String, String> cookies,
Map<String, Object> 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<String, String> getCookies() {
return cookies;
}
public Map<String, Object> 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);
}
}

View File

@@ -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();
}
}

View File

@@ -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<String, ResponseCookie> cookies;
private final RequestData requestData;
public ResponseData(HttpStatus httpStatus, HttpHeaders headers, MultiValueMap<String, ResponseCookie> 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<String, ResponseCookie> 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);
}
}

View File

@@ -90,7 +90,7 @@ public class RetryLoadBalancerInterceptor implements ClientHttpRequestIntercepto
Set<LoadBalancerLifecycle> 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<RetryableRequestContext> 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()) {

View File

@@ -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();
}
}

View File

@@ -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<LoadBalancerLifecycle> 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<ClientRequestContext> lbRequest = new DefaultRequest<>(
new ClientRequestContext(clientRequest, hint));
RequestData requestData = new RequestData(clientRequest);
DefaultRequest<RequestDataContext> 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<ClientResponse, ServiceInstance>(
lifecycle -> lifecycle.onComplete(new CompletionContext<ResponseData, ServiceInstance>(
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<Response<ServiceInstance>> choose(String serviceId, Request<ClientRequestContext> request) {
protected Mono<Response<ServiceInstance>> choose(String serviceId, Request<RequestDataContext> request) {
ReactiveLoadBalancer<ServiceInstance> loadBalancer = loadBalancerFactory.getInstance(serviceId);
if (loadBalancer == null) {
return Mono.just(new EmptyResponse());

View File

@@ -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<LoadBalancerLifecycle> 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<RetryableRequestContext> 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)) {

View File

@@ -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<Request<Object>> 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" })

View File

@@ -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

View File

@@ -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<List<ServiceInstance>> 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<String, String> 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();
}

View File

@@ -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<ClientRequestContext> request = new DefaultRequest<>(new ClientRequestContext(clientRequest));
Request<RequestDataContext> request = new DefaultRequest<>(
new RequestDataContext(new RequestData(clientRequest)));
List<ServiceInstance> 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<ClientRequestContext> request = new DefaultRequest<>(new ClientRequestContext(clientRequest));
Request<RequestDataContext> request = new DefaultRequest<>(
new RequestDataContext(new RequestData(clientRequest)));
List<ServiceInstance> serviceInstances = supplier.get(request).blockFirst();
@@ -100,48 +99,8 @@ class RequestBasedStickySessionServiceInstanceListSupplierTests {
@Test
void shouldReturnAllInstancesFromDelegateIfClientRequestHasNoCookie() {
when(clientRequest.cookies()).thenReturn(new HttpHeaders());
Request<ClientRequestContext> request = new DefaultRequest<>(new ClientRequestContext(clientRequest));
List<ServiceInstance> serviceInstances = supplier.get(request).blockFirst();
assertThat(serviceInstances).hasSize(3);
}
@Test
void shouldReturnInstanceBasedOnCookieFromServerHttpRequest() {
MultiValueMap<String, HttpCookie> cookies = new LinkedMultiValueMap<>();
cookies.add(properties.getStickySession().getInstanceIdCookieName(),
new HttpCookie(properties.getStickySession().getInstanceIdCookieName(), "test-1"));
when(serverHttpRequest.getCookies()).thenReturn(cookies);
Request<ServerHttpRequestContext> request = new DefaultRequest<>(
new ServerHttpRequestContext(serverHttpRequest));
List<ServiceInstance> serviceInstances = supplier.get(request).blockFirst();
assertThat(serviceInstances).hasSize(1);
assertThat(serviceInstances.get(0).getInstanceId()).isEqualTo("test-1");
}
@Test
void shouldReturnAllDelegateInstancesIfInstanceBasedOnCookieFromServerHttpRequestNotFound() {
MultiValueMap<String, HttpCookie> cookies = new LinkedMultiValueMap<>();
cookies.add(properties.getStickySession().getInstanceIdCookieName(),
new HttpCookie(properties.getStickySession().getInstanceIdCookieName(), "test-4"));
when(serverHttpRequest.getCookies()).thenReturn(cookies);
Request<ServerHttpRequestContext> request = new DefaultRequest<>(
new ServerHttpRequestContext(serverHttpRequest));
List<ServiceInstance> serviceInstances = supplier.get(request).blockFirst();
assertThat(serviceInstances).hasSize(3);
}
@Test
void shouldReturnAllInstancesFromDelegateIfServerHttpRequestHasNoCookie() {
MultiValueMap<String, HttpCookie> cookies = new LinkedMultiValueMap<>();
when(serverHttpRequest.getCookies()).thenReturn(cookies);
Request<ServerHttpRequestContext> request = new DefaultRequest<>(
new ServerHttpRequestContext(serverHttpRequest));
Request<RequestDataContext> request = new DefaultRequest<>(
new RequestDataContext(new RequestData(clientRequest)));
List<ServiceInstance> serviceInstances = supplier.get(request).blockFirst();