diff --git a/spring-cloud-gateway-core/pom.xml b/spring-cloud-gateway-core/pom.xml index 8671eb16..6dcd5ec4 100644 --- a/spring-cloud-gateway-core/pom.xml +++ b/spring-cloud-gateway-core/pom.xml @@ -47,6 +47,11 @@ spring-cloud-starter-netflix-hystrix true + + org.springframework.cloud + spring-cloud-loadbalancer + true + org.springframework.boot spring-boot-devtools diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayLoadBalancerClientAutoConfiguration.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayLoadBalancerClientAutoConfiguration.java index 119c3adc..ef586195 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayLoadBalancerClientAutoConfiguration.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayLoadBalancerClientAutoConfiguration.java @@ -36,8 +36,6 @@ import org.springframework.web.reactive.DispatcherHandler; @AutoConfigureAfter(RibbonAutoConfiguration.class) public class GatewayLoadBalancerClientAutoConfiguration { - // GlobalFilter beans - @Bean @ConditionalOnBean(LoadBalancerClient.class) @ConditionalOnMissingBean(LoadBalancerClientFilter.class) diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java new file mode 100644 index 00000000..f17ff32c --- /dev/null +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/config/GatewayReactiveLoadBalancerClientAutoConfiguration.java @@ -0,0 +1,48 @@ +/* + * Copyright 2013-2018 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 + * + * http://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.config; + +import org.springframework.boot.autoconfigure.AutoConfigureAfter; +import org.springframework.boot.autoconfigure.AutoConfigureBefore; +import org.springframework.boot.autoconfigure.condition.ConditionalOnBean; +import org.springframework.boot.autoconfigure.condition.ConditionalOnClass; +import org.springframework.boot.autoconfigure.condition.ConditionalOnMissingBean; +import org.springframework.cloud.client.loadbalancer.reactive.ReactiveLoadBalancer; +import org.springframework.cloud.gateway.filter.ReactiveLoadBalancerClientFilter; +import org.springframework.cloud.loadbalancer.config.LoadBalancerAutoConfiguration; +import org.springframework.cloud.loadbalancer.support.LoadBalancerClientFactory; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.web.reactive.DispatcherHandler; + +/** + * @author Spencer Gibb + */ +@Configuration +@ConditionalOnClass({ReactiveLoadBalancer.class, LoadBalancerAutoConfiguration.class, DispatcherHandler.class}) +@AutoConfigureBefore(GatewayLoadBalancerClientAutoConfiguration.class) +@AutoConfigureAfter(LoadBalancerAutoConfiguration.class) +public class GatewayReactiveLoadBalancerClientAutoConfiguration { + + @Bean + @ConditionalOnBean(LoadBalancerClientFactory.class) + @ConditionalOnMissingBean(ReactiveLoadBalancerClientFilter.class) + public ReactiveLoadBalancerClientFilter loadBalancerClientFilter(LoadBalancerClientFactory clientFactory) { + return new ReactiveLoadBalancerClientFilter(clientFactory); + } +} diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilter.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilter.java index ff12c753..4fe7fb52 100644 --- a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilter.java +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilter.java @@ -18,12 +18,14 @@ package org.springframework.cloud.gateway.filter; import java.net.URI; -import java.util.Map; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import reactor.core.publisher.Mono; + import org.springframework.cloud.client.ServiceInstance; import org.springframework.cloud.client.loadbalancer.LoadBalancerClient; +import org.springframework.cloud.gateway.support.DelegatingServiceInstance; import org.springframework.cloud.gateway.support.NotFoundException; import org.springframework.core.Ordered; import org.springframework.web.server.ServerWebExchange; @@ -32,8 +34,6 @@ import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.G import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_SCHEME_PREFIX_ATTR; import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.addOriginalRequestUrl; -import reactor.core.publisher.Mono; - /** * @author Spencer Gibb * @author Tim Ysewyn @@ -92,53 +92,4 @@ public class LoadBalancerClientFilter implements GlobalFilter, Ordered { return loadBalancer.choose(((URI) exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR)).getHost()); } - class DelegatingServiceInstance implements ServiceInstance { - final ServiceInstance delegate; - private String overrideScheme; - - DelegatingServiceInstance(ServiceInstance delegate, String overrideScheme) { - this.delegate = delegate; - this.overrideScheme = overrideScheme; - } - - @Override - public String getServiceId() { - return delegate.getServiceId(); - } - - @Override - public String getHost() { - return delegate.getHost(); - } - - @Override - public int getPort() { - return delegate.getPort(); - } - - @Override - public boolean isSecure() { - return delegate.isSecure(); - } - - @Override - public URI getUri() { - return delegate.getUri(); - } - - @Override - public Map getMetadata() { - return delegate.getMetadata(); - } - - @Override - public String getScheme() { - String scheme = delegate.getScheme(); - if (scheme != null) { - return scheme; - } - return this.overrideScheme; - } - - } } diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java new file mode 100644 index 00000000..7f6b6714 --- /dev/null +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/filter/ReactiveLoadBalancerClientFilter.java @@ -0,0 +1,166 @@ +/* + * Copyright 2013-2018 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 + * + * http://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; + +import java.net.URI; +import java.util.HashMap; +import java.util.Map; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; +import org.springframework.cloud.client.loadbalancer.reactive.ReactiveLoadBalancer; +import org.springframework.cloud.client.loadbalancer.reactive.Request; +import org.springframework.cloud.client.loadbalancer.reactive.Response; +import org.springframework.cloud.gateway.support.DelegatingServiceInstance; +import org.springframework.cloud.loadbalancer.core.ReactorLoadBalancer; +import org.springframework.cloud.loadbalancer.support.LoadBalancerClientFactory; +import org.springframework.util.StringUtils; +import org.springframework.web.util.UriComponentsBuilder; +import reactor.core.publisher.Mono; + +import org.springframework.cloud.client.ServiceInstance; +import org.springframework.cloud.gateway.support.NotFoundException; +import org.springframework.core.Ordered; +import org.springframework.web.server.ServerWebExchange; + +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_SCHEME_PREFIX_ATTR; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.addOriginalRequestUrl; + +/** + * @author Spencer Gibb + * @author Tim Ysewyn + */ +public class ReactiveLoadBalancerClientFilter implements GlobalFilter, Ordered { + + private static final Log log = LogFactory.getLog(ReactiveLoadBalancerClientFilter.class); + public static final int LOAD_BALANCER_CLIENT_FILTER_ORDER = 10150; + + protected final LoadBalancerClientFactory clientFactory; + + public ReactiveLoadBalancerClientFilter(LoadBalancerClientFactory clientFactory) { + this.clientFactory = clientFactory; + } + + @Override + public int getOrder() { + return LOAD_BALANCER_CLIENT_FILTER_ORDER; + } + + @Override + public Mono filter(ServerWebExchange exchange, GatewayFilterChain chain) { + URI url = exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR); + String schemePrefix = exchange.getAttribute(GATEWAY_SCHEME_PREFIX_ATTR); + if (url == null || (!"lb".equals(url.getScheme()) && !"lb".equals(schemePrefix))) { + return chain.filter(exchange); + } + //preserve the original url + addOriginalRequestUrl(exchange, url); + + if (log.isTraceEnabled()) { + log.trace("ReactiveLoadBalancerClientFilter url before: " + url); + } + + return choose(exchange).doOnNext(response -> { + + if (!response.hasServer()) { + throw new NotFoundException("Unable to find instance for " + url.getHost()); + } + + URI uri = exchange.getRequest().getURI(); + + // if the `lb:` mechanism was used, use `` as the default, + // if the loadbalancer doesn't provide one. + String overrideScheme = null; + if (schemePrefix != null) { + overrideScheme = url.getScheme(); + } + + DelegatingServiceInstance serviceInstance = new DelegatingServiceInstance(response.getServer(), overrideScheme); + // URI requestUrl = clientFactory.reconstructURI(serviceInstance, uri); + URI requestUrl = updateUri(uri, serviceInstance); + + if (log.isTraceEnabled()) { + log.trace("LoadBalancerClientFilter url chosen: " + requestUrl); + } + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, requestUrl); + }).then(chain.filter(exchange)); + } + + protected Mono> choose(ServerWebExchange exchange) { + URI uri = exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR); + ReactorLoadBalancer loadBalancer = this.clientFactory.getInstance(uri.getHost(), ReactorLoadBalancer.class, ServiceInstance.class); + if (loadBalancer == null) { + throw new NotFoundException("No loadbalancer available for " + uri.getHost()); + } + return loadBalancer.choose(createRequest()); + } + + protected Request createRequest() { + return ReactiveLoadBalancer.REQUEST; + } + + + private static final Map unsecureSchemeMapping; + static + { + unsecureSchemeMapping = new HashMap<>(); + unsecureSchemeMapping.put("http", "https"); + unsecureSchemeMapping.put("ws", "wss"); + } + + /** + * Replace the scheme to the secure variant if needed. If the {@link #unsecureSchemeMapping} map contains the uri + * scheme and {@link ServiceInstance#isSecure()} is true, update the scheme. + * This assumes the uri is already encoded to avoid double encoding. + * + * @param uri + * @param serviceInstance + * @return + */ + static String updateToSecureScheme(URI uri, ServiceInstance serviceInstance) { + String scheme = uri.getScheme(); + + if (StringUtils.isEmpty(scheme)) { + scheme = "http"; + } + + if (!StringUtils.isEmpty(uri.toString()) + && unsecureSchemeMapping.containsKey(scheme) + && serviceInstance.isSecure()) { + return unsecureSchemeMapping.get(scheme); + } + return scheme; + } + + static URI updateUri(URI uri, ServiceInstance serviceInstance) { + UriComponentsBuilder builder = UriComponentsBuilder + .fromUri(uri) + .scheme(updateToSecureScheme(uri, serviceInstance)) + .host(serviceInstance.getHost()) + .port(serviceInstance.getPort()); + // follow up with https://jira.spring.io/browse/SPR-17039 + if (uri.getRawQuery() != null) { + // When building the URI, UriComponentsBuilder verify the allowed characters and does not + // support the '+' so we replace it for its equivalent '%20'. + // See issue https://jira.spring.io/browse/SPR-10172 + builder.replaceQuery(uri.getRawQuery().replace("+", "%20")); + } + return builder.build(true).toUri(); + } +} diff --git a/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/DelegatingServiceInstance.java b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/DelegatingServiceInstance.java new file mode 100644 index 00000000..87e4cfa3 --- /dev/null +++ b/spring-cloud-gateway-core/src/main/java/org/springframework/cloud/gateway/support/DelegatingServiceInstance.java @@ -0,0 +1,77 @@ +/* + * Copyright 2013-2018 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 + * + * http://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.support; + +import org.springframework.cloud.client.ServiceInstance; + +import java.net.URI; +import java.util.Map; + +public class DelegatingServiceInstance implements ServiceInstance { + final ServiceInstance delegate; + private String overrideScheme; + + public DelegatingServiceInstance(ServiceInstance delegate, String overrideScheme) { + this.delegate = delegate; + this.overrideScheme = overrideScheme; + } + + @Override + public String getServiceId() { + return delegate.getServiceId(); + } + + @Override + public String getHost() { + return delegate.getHost(); + } + + @Override + public int getPort() { + return delegate.getPort(); + } + + @Override + public boolean isSecure() { + //TODO: move to map + if ("https".equals(this.overrideScheme) || "wss".equals(this.overrideScheme)) { + return true; + } + return delegate.isSecure(); + } + + @Override + public URI getUri() { + return delegate.getUri(); + } + + @Override + public Map getMetadata() { + return delegate.getMetadata(); + } + + @Override + public String getScheme() { + String scheme = delegate.getScheme(); + if (scheme != null) { + return scheme; + } + return this.overrideScheme; + } + +} diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ReactorLoadBalancerClientFilterTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ReactorLoadBalancerClientFilterTests.java new file mode 100644 index 00000000..0cbb503d --- /dev/null +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/ReactorLoadBalancerClientFilterTests.java @@ -0,0 +1,307 @@ +/* + * Copyright 2013-2018 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 + * + * http://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; + +import java.net.URI; +import java.util.LinkedHashSet; + +import org.junit.Before; +import org.junit.Ignore; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.mockito.ArgumentCaptor; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.junit.MockitoJUnitRunner; +import reactor.core.publisher.Mono; + +import org.springframework.cloud.client.DefaultServiceInstance; +import org.springframework.cloud.client.ServiceInstance; +import org.springframework.cloud.gateway.support.NotFoundException; +import org.springframework.cloud.loadbalancer.core.ReactorLoadBalancer; +import org.springframework.cloud.loadbalancer.core.RoundRobinLoadBalancer; +import org.springframework.cloud.loadbalancer.support.LoadBalancerClientFactory; +import org.springframework.cloud.loadbalancer.support.ServiceInstanceSuppliers; +import org.springframework.http.HttpMethod; +import org.springframework.mock.http.server.reactive.MockServerHttpRequest; +import org.springframework.mock.web.server.MockServerWebExchange; +import org.springframework.web.server.ServerWebExchange; +import org.springframework.web.util.UriComponentsBuilder; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoMoreInteractions; +import static org.mockito.Mockito.verifyZeroInteractions; +import static org.mockito.Mockito.when; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_ORIGINAL_REQUEST_URL_ATTR; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR; +import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_SCHEME_PREFIX_ATTR; + +/** + * @author Spencer Gibb + * @author Tim Ysewyn + */ +@RunWith(MockitoJUnitRunner.class) +public class ReactorLoadBalancerClientFilterTests { + + private ServerWebExchange exchange; + + @Mock + private GatewayFilterChain chain; + + @Mock + private LoadBalancerClientFactory clientFactory; + + @InjectMocks + private ReactiveLoadBalancerClientFilter filter; + + @Before + public void setup() { + exchange = MockServerWebExchange.from(MockServerHttpRequest.get("/mypath").build()); + } + + @Test + public void shouldNotFilterWhenGatewayRequestUrlIsMissing() { + filter.filter(exchange, chain); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + verifyZeroInteractions(clientFactory); + } + + @Test + public void shouldNotFilterWhenGatewayRequestUrlSchemeIsNotLb() { + URI uri = UriComponentsBuilder.fromUriString("http://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + + filter.filter(exchange, chain); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + verifyZeroInteractions(clientFactory); + } + + @Test(expected = NotFoundException.class) + public void shouldThrowExceptionWhenNoServiceInstanceIsFound() { + URI uri = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + + filter.filter(exchange, chain).block(); + } + + @Test + public void shouldFilter() { + URI url = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, url); + + ServiceInstance serviceInstance = new DefaultServiceInstance("myservice", "localhost", 8080, true); + + when(clientFactory.getInstance("myservice", ReactorLoadBalancer.class, ServiceInstance.class)) + .thenReturn(new RoundRobinLoadBalancer("myservice", + ServiceInstanceSuppliers.toProvider("myservice", serviceInstance), + -1)); + + when(chain.filter(exchange)).thenReturn(Mono.empty()); + + filter.filter(exchange, chain).block(); + + assertThat((LinkedHashSet)exchange.getAttribute(GATEWAY_ORIGINAL_REQUEST_URL_ATTR)).contains(url); + + verify(clientFactory).getInstance("myservice", ReactorLoadBalancer.class, ServiceInstance.class); + + verifyNoMoreInteractions(clientFactory); + + assertThat((URI)exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR)).isEqualTo(URI.create("https://localhost:8080/mypath")); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + } + + @Test + public void happyPath() { + MockServerHttpRequest request = MockServerHttpRequest + .get("http://localhost/get?a=b") + .build(); + + URI lbUri = URI.create("lb://service1?a=b"); + ServerWebExchange webExchange = testFilter(request, lbUri); + URI uri = webExchange.getRequiredAttribute(GATEWAY_REQUEST_URL_ATTR); + assertThat(uri).hasScheme("http").hasHost("service1-host1") + .hasParameter("a", "b"); + } + + @Test + public void noQueryParams() { + MockServerHttpRequest request = MockServerHttpRequest + .get("http://localhost/get") + .build(); + + ServerWebExchange webExchange = testFilter(request, URI.create("lb://service1")); + URI uri = webExchange.getRequiredAttribute(GATEWAY_REQUEST_URL_ATTR); + assertThat(uri).hasScheme("http").hasHost("service1-host1"); + } + + @Test + public void encodedParameters() { + URI url = UriComponentsBuilder.fromUriString("http://localhost/get?a=b&c=d[]").buildAndExpand().encode().toUri(); + + MockServerHttpRequest request = MockServerHttpRequest + .method(HttpMethod.GET, url) + .build(); + + URI lbUrl = UriComponentsBuilder.fromUriString("lb://service1?a=b&c=d[]").buildAndExpand().encode().toUri(); + + // prove that it is encoded + assertThat(lbUrl.getRawQuery()).isEqualTo("a=b&c=d%5B%5D"); + + assertThat(lbUrl).hasParameter("c", "d[]"); + + ServerWebExchange webExchange = testFilter(request, lbUrl); + URI uri = webExchange.getRequiredAttribute(GATEWAY_REQUEST_URL_ATTR); + assertThat(uri).hasScheme("http").hasHost("service1-host1") + .hasParameter("a", "b") + .hasParameter("c", "d[]"); + + // prove that it is not double encoded + assertThat(uri.getRawQuery()).isEqualTo("a=b&c=d%5B%5D"); + } + + @Ignore //FIXME: 2.1.0 + @Test + public void unencodedParameters() { + URI url = URI.create("http://localhost/get?a=b&c=d[]"); + + MockServerHttpRequest request = MockServerHttpRequest + .method(HttpMethod.GET, url) + .build(); + + URI lbUrl = URI.create("lb://service1?a=b&c=d[]"); + + // prove that it is unencoded + assertThat(lbUrl.getRawQuery()).isEqualTo("a=b&c=d[]"); + + ServerWebExchange webExchange = testFilter(request, lbUrl); + + URI uri = webExchange.getRequiredAttribute(GATEWAY_REQUEST_URL_ATTR); + assertThat(uri).hasScheme("http").hasHost("service1-host1") + .hasParameter("a", "b") + .hasParameter("c", "d[]"); + + // prove that it is NOT encoded + assertThat(uri.getRawQuery()).isEqualTo("a=b&c=d[]"); + } + + @Test + public void happyPathWithAttributeRatherThanScheme() { + MockServerHttpRequest request = MockServerHttpRequest + .get("ws://localhost/get?a=b") + .build(); + + URI lbUri = URI.create("ws://service1?a=b"); + + exchange = MockServerWebExchange.from(request); + exchange.getAttributes().put(GATEWAY_SCHEME_PREFIX_ATTR, "lb"); + + ServerWebExchange webExchange = testFilter(exchange, lbUri); + URI uri = webExchange.getRequiredAttribute(GATEWAY_REQUEST_URL_ATTR); + assertThat(uri).hasScheme("ws").hasHost("service1-host1") + .hasParameter("a", "b"); + } + + @Test + public void shouldNotFilterWhenGatewaySchemePrefixAttrIsNotLb() { + URI uri = UriComponentsBuilder.fromUriString("http://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + exchange.getAttributes().put(GATEWAY_SCHEME_PREFIX_ATTR, "xx"); + + filter.filter(exchange, chain); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + verifyZeroInteractions(clientFactory); + } + + @Ignore //FIXME: 2.1.0 + @Test + public void shouldSelectSpecifiedServer() { + /*URI uri1 = UriComponentsBuilder.fromUriString("lb://myservice").port(11111).build().toUri(); + URI uri2 = UriComponentsBuilder.fromUriString("lb://myservice").port(22222).build().toUri(); + + SpringClientFactory clientFactory = mock(SpringClientFactory.class); + ILoadBalancer loadBalancer = mock(ILoadBalancer.class); + when(clientFactory.getLoadBalancerContext("myservice")).thenReturn(new RibbonLoadBalancerContext(loadBalancer)); + when(clientFactory.getLoadBalancer("myservice")).thenReturn(loadBalancer); + + when(loadBalancer.chooseServer("11111")).thenReturn(new Server("myservice-host1", 8081)); + when(loadBalancer.chooseServer("22222")).thenReturn(new Server("myservice-host2", 8081)); + + LoadBalancerClient loadBalancerClient = new RibbonLoadBalancerClient(clientFactory) { + private String loadBalancerKey; + public ServiceInstance choose(String serviceId) { + String[] strings = serviceId.split("<<>>"); + loadBalancerKey = strings[1]; + return super.choose(strings[0]); + } + protected Server getServer(ILoadBalancer loadBalancer) { + return loadBalancer == null ? null : loadBalancer.chooseServer(StringUtils.isEmpty(loadBalancerKey) ? "default" : loadBalancerKey); + } + }; + + LoadBalancerClientFilter loadBalancerClientFilter = new LoadBalancerClientFilter(loadBalancerClient) { + protected ServiceInstance choose(ServerWebExchange exchange) { + URI attribute = (URI) exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR); + return loadBalancer.choose(attribute.getHost() + "<<>>" + attribute.getPort()); + } + }; + + MockServerHttpRequest request = MockServerHttpRequest + .get("http://localhost/get") + .build(); + ServerWebExchange exchange = MockServerWebExchange.from(request); + + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri1); + loadBalancerClientFilter.filter(exchange, chain); + assertThat(((URI)exchange.getAttributes().get(GATEWAY_REQUEST_URL_ATTR)).getHost()).isEqualTo("myservice-host1"); + + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri2); + loadBalancerClientFilter.filter(exchange, chain); + assertThat(((URI)exchange.getAttributes().get(GATEWAY_REQUEST_URL_ATTR)).getHost()).isEqualTo("myservice-host2");*/ + } + + private ServerWebExchange testFilter(MockServerHttpRequest request, URI uri) { + return testFilter(MockServerWebExchange.from(request), uri); + } + + private ServerWebExchange testFilter(ServerWebExchange exchange, URI uri) { + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + + ArgumentCaptor captor = ArgumentCaptor.forClass(ServerWebExchange.class); + when(chain.filter(captor.capture())).thenReturn(Mono.empty()); + + when(clientFactory.getInstance("service1", ReactorLoadBalancer.class, ServiceInstance.class)) + .thenReturn(new RoundRobinLoadBalancer("service1", + ServiceInstanceSuppliers.toProvider("service1", + new DefaultServiceInstance("service1", "service1-host1", 8081, false)), + -1)); + + ReactiveLoadBalancerClientFilter filter = new ReactiveLoadBalancerClientFilter(clientFactory); + filter.filter(exchange, chain).block(); + + return captor.getValue(); + } +}