diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java index 62e98fa4..4ac6890d 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java @@ -20,9 +20,7 @@ import org.springframework.boot.autoconfigure.AutoConfigureAfter; import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; -import org.springframework.boot.autoconfigure.condition.ConditionalOnProperty; import org.springframework.boot.context.properties.EnableConfigurationProperties; -import org.springframework.cloud.client.loadbalancer.LoadBalancerProperties; import org.springframework.cloud.client.loadbalancer.reactive.ReactiveLoadBalancer; import org.springframework.cloud.gateway.config.conditional.ConditionalOnEnabledGlobalFilter; import org.springframework.cloud.gateway.filter.LoadBalancerServiceInstanceCookieFilter; @@ -50,19 +48,17 @@ public class GatewayReactiveLoadBalancerClientAutoConfiguration { @ConditionalOnMissingBean(ReactiveLoadBalancerClientFilter.class) @ConditionalOnEnabledGlobalFilter public ReactiveLoadBalancerClientFilter gatewayLoadBalancerClientFilter(LoadBalancerClientFactory clientFactory, - GatewayLoadBalancerProperties properties, LoadBalancerProperties loadBalancerProperties) { - return new ReactiveLoadBalancerClientFilter(clientFactory, properties, loadBalancerProperties); + GatewayLoadBalancerProperties properties) { + return new ReactiveLoadBalancerClientFilter(clientFactory, properties); } @Bean - @ConditionalOnBean(ReactiveLoadBalancerClientFilter.class) - @ConditionalOnProperty(value = "spring.cloud.loadbalancer.sticky-session.add-service-instance-cookie", - havingValue = "true") + @ConditionalOnBean({ ReactiveLoadBalancerClientFilter.class, LoadBalancerClientFactory.class }) @ConditionalOnMissingBean @ConditionalOnEnabledGlobalFilter public LoadBalancerServiceInstanceCookieFilter loadBalancerServiceInstanceCookieFilter( - LoadBalancerProperties loadBalancerProperties) { - return new LoadBalancerServiceInstanceCookieFilter(loadBalancerProperties); + LoadBalancerClientFactory loadBalancerClientFactory) { + return new LoadBalancerServiceInstanceCookieFilter(loadBalancerClientFactory); } } diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilter.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilter.java index 9581a9d2..6f05c39f 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilter.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilter.java @@ -24,6 +24,7 @@ import reactor.core.publisher.Mono; import org.springframework.cloud.client.ServiceInstance; import org.springframework.cloud.client.loadbalancer.LoadBalancerProperties; import org.springframework.cloud.client.loadbalancer.Response; +import org.springframework.cloud.client.loadbalancer.reactive.ReactiveLoadBalancer; import org.springframework.core.Ordered; import org.springframework.http.HttpCookie; import org.springframework.http.HttpHeaders; @@ -43,19 +44,37 @@ import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.G */ public class LoadBalancerServiceInstanceCookieFilter implements GlobalFilter, Ordered { - private final LoadBalancerProperties loadBalancerProperties; + private LoadBalancerProperties loadBalancerProperties; + private ReactiveLoadBalancer.Factory loadBalancerClientFactory; + + /** + * @deprecated in favour of + * {@link LoadBalancerServiceInstanceCookieFilter#LoadBalancerServiceInstanceCookieFilter(ReactiveLoadBalancer.Factory)} + */ + @Deprecated public LoadBalancerServiceInstanceCookieFilter(LoadBalancerProperties loadBalancerProperties) { this.loadBalancerProperties = loadBalancerProperties; } + public LoadBalancerServiceInstanceCookieFilter( + ReactiveLoadBalancer.Factory loadBalancerClientFactory) { + this.loadBalancerClientFactory = loadBalancerClientFactory; + } + @Override public Mono filter(ServerWebExchange exchange, GatewayFilterChain chain) { Response serviceInstanceResponse = exchange.getAttribute(GATEWAY_LOADBALANCER_RESPONSE_ATTR); if (serviceInstanceResponse == null || !serviceInstanceResponse.hasServer()) { return chain.filter(exchange); } - String instanceIdCookieName = loadBalancerProperties.getStickySession().getInstanceIdCookieName(); + LoadBalancerProperties properties = loadBalancerClientFactory != null + ? loadBalancerClientFactory.getProperties(serviceInstanceResponse.getServer().getServiceId()) + : loadBalancerProperties; + if (!properties.getStickySession().isAddServiceInstanceCookie()) { + return chain.filter(exchange); + } + String instanceIdCookieName = properties.getStickySession().getInstanceIdCookieName(); if (!StringUtils.hasText(instanceIdCookieName)) { return chain.filter(exchange); } diff --git a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java index 725859f9..3de22769 100644 --- a/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java +++ b/spring-cloud-gateway-server/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java @@ -72,13 +72,21 @@ public class ReactiveLoadBalancerClientFilter implements GlobalFilter, Ordered { private final GatewayLoadBalancerProperties properties; - private final LoadBalancerProperties loadBalancerProperties; - + /** + * @deprecated in favour of + * {@link ReactiveLoadBalancerClientFilter#ReactiveLoadBalancerClientFilter(LoadBalancerClientFactory, GatewayLoadBalancerProperties)} + */ + @Deprecated public ReactiveLoadBalancerClientFilter(LoadBalancerClientFactory clientFactory, GatewayLoadBalancerProperties properties, LoadBalancerProperties loadBalancerProperties) { this.clientFactory = clientFactory; this.properties = properties; - this.loadBalancerProperties = loadBalancerProperties; + } + + public ReactiveLoadBalancerClientFilter(LoadBalancerClientFactory clientFactory, + GatewayLoadBalancerProperties properties) { + this.clientFactory = clientFactory; + this.properties = properties; } @Override @@ -105,8 +113,8 @@ public class ReactiveLoadBalancerClientFilter implements GlobalFilter, Ordered { Set supportedLifecycleProcessors = LoadBalancerLifecycleValidator .getSupportedLifecycleProcessors(clientFactory.getInstances(serviceId, LoadBalancerLifecycle.class), RequestDataContext.class, ResponseData.class, ServiceInstance.class); - DefaultRequest lbRequest = new DefaultRequest<>(new RequestDataContext( - new RequestData(exchange.getRequest()), getHint(serviceId, loadBalancerProperties.getHint()))); + DefaultRequest lbRequest = new DefaultRequest<>( + new RequestDataContext(new RequestData(exchange.getRequest()), getHint(serviceId))); return choose(lbRequest, serviceId, supportedLifecycleProcessors).doOnNext(response -> { if (!response.hasServer()) { @@ -164,7 +172,9 @@ public class ReactiveLoadBalancerClientFilter implements GlobalFilter, Ordered { return loadBalancer.choose(lbRequest); } - private String getHint(String serviceId, Map hints) { + private String getHint(String serviceId) { + LoadBalancerProperties loadBalancerProperties = clientFactory.getProperties(serviceId); + Map hints = loadBalancerProperties.getHint(); String defaultHint = hints.getOrDefault("default", "default"); String hintPropertyValue = hints.get(serviceId); return hintPropertyValue != null ? hintPropertyValue : defaultHint; diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilterTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilterTests.java index 9cfad225..2821afca 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilterTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerServiceInstanceCookieFilterTests.java @@ -16,6 +16,7 @@ package org.springframework.cloud.gateway.filter; +import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.mockito.ArgumentCaptor; import reactor.core.publisher.Mono; @@ -54,6 +55,11 @@ class LoadBalancerServiceInstanceCookieFilterTests { private final LoadBalancerServiceInstanceCookieFilter filter = new LoadBalancerServiceInstanceCookieFilter( properties); + @BeforeEach + void setUp() { + properties.getStickySession().setAddServiceInstanceCookie(true); + } + @Test void shouldAddServiceInstanceCookieHeader() { exchange.getAttributes().put(GATEWAY_LOADBALANCER_RESPONSE_ATTR, diff --git a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilterTests.java b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilterTests.java index 7d31eb71..6c08b199 100644 --- a/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilterTests.java +++ b/spring-cloud-gateway-server/src/test/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilterTests.java @@ -124,6 +124,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void shouldThrowExceptionWhenNoServiceInstanceIsFound() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); assertThatExceptionOfType(NotFoundException.class).isThrownBy(() -> { URI uri = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); @@ -135,6 +136,7 @@ class ReactiveLoadBalancerClientFilterTests { @SuppressWarnings("unchecked") @Test void shouldFilter() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); URI url = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, url); @@ -166,6 +168,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void happyPath() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); MockServerHttpRequest request = MockServerHttpRequest.get("http://localhost/get?a=b").build(); URI lbUri = URI.create("lb://service1?a=b"); @@ -176,6 +179,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void noQueryParams() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); MockServerHttpRequest request = MockServerHttpRequest.get("http://localhost/get").build(); ServerWebExchange webExchange = testFilter(request, URI.create("lb://service1")); @@ -185,6 +189,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void encodedParameters() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); URI url = UriComponentsBuilder.fromUriString("http://localhost/get?a=b&c=d[]").buildAndExpand().encode() .toUri(); @@ -207,6 +212,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void unencodedParameters() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); URI url = URI.create("http://localhost/get?a=b&c=d[]"); MockServerHttpRequest request = MockServerHttpRequest.method(HttpMethod.GET, url).build(); @@ -227,6 +233,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void happyPathWithAttributeRatherThanScheme() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); MockServerHttpRequest request = MockServerHttpRequest.get("ws://localhost/get?a=b").build(); URI lbUri = URI.create("ws://service1?a=b"); @@ -254,6 +261,7 @@ class ReactiveLoadBalancerClientFilterTests { @Test void shouldThrow4O4ExceptionWhenNoServiceInstanceIsFound() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); URI uri = UriComponentsBuilder.fromUriString("lb://service1").build().toUri(); exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); RoundRobinLoadBalancer loadBalancer = new RoundRobinLoadBalancer( @@ -274,6 +282,8 @@ class ReactiveLoadBalancerClientFilterTests { @SuppressWarnings("unchecked") @Test void shouldOverrideSchemeUsingIsSecure() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); URI url = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); ServerWebExchange exchange = MockServerWebExchange .from(MockServerHttpRequest.get("https://localhost:9999/mypath").build()); @@ -299,6 +309,7 @@ class ReactiveLoadBalancerClientFilterTests { void shouldPassRequestToLoadBalancer() { String hint = "test"; when(loadBalancerProperties.getHint()).thenReturn(buildHints(hint)); + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); MockServerHttpRequest request = MockServerHttpRequest.get("http://localhost/get?a=b").build(); URI lbUri = URI.create("lb://service1?a=b"); ServerWebExchange serverWebExchange = mock(ServerWebExchange.class); @@ -323,6 +334,7 @@ class ReactiveLoadBalancerClientFilterTests { @SuppressWarnings({ "unchecked", "rawtypes" }) @Test void loadBalancerLifecycleCallbacksExecutedForSuccess() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); LoadBalancerLifecycle lifecycleProcessor = mock(LoadBalancerLifecycle.class); ServiceInstance serviceInstance = new DefaultServiceInstance("myservice1", "myservice", "localhost", 8080, false); @@ -342,6 +354,7 @@ class ReactiveLoadBalancerClientFilterTests { @SuppressWarnings({ "unchecked", "rawtypes" }) @Test void loadBalancerLifecycleCallbacksExecutedForDiscard() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); LoadBalancerLifecycle lifecycleProcessor = mock(LoadBalancerLifecycle.class); ServiceInstance serviceInstance = null; ServerWebExchange serverWebExchange = mockExchange(serviceInstance, lifecycleProcessor, false); @@ -358,6 +371,7 @@ class ReactiveLoadBalancerClientFilterTests { @SuppressWarnings({ "unchecked", "rawtypes" }) @Test void loadBalancerLifecycleCallbacksExecutedForFailed() { + when(clientFactory.getProperties(any())).thenReturn(loadBalancerProperties); LoadBalancerLifecycle lifecycleProcessor = mock(LoadBalancerLifecycle.class); ServiceInstance serviceInstance = new DefaultServiceInstance("myservice1", "myservice", "localhost", 8080, false);