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 05c43bf3..f1f49909 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 @@ -26,16 +26,15 @@ import org.springframework.cloud.client.loadbalancer.LoadBalancerClient; import org.springframework.cloud.gateway.support.NotFoundException; import org.springframework.core.Ordered; import org.springframework.web.server.ServerWebExchange; -import org.springframework.web.util.UriComponentsBuilder; import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.GATEWAY_REQUEST_URL_ATTR; import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.addOriginalRequestUrl; -import static org.springframework.cloud.gateway.support.ServerWebExchangeUtils.containsEncodedQuery; import reactor.core.publisher.Mono; /** * @author Spencer Gibb + * @author Tim Ysewyn */ public class LoadBalancerClientFilter implements GlobalFilter, Ordered { @@ -70,15 +69,9 @@ public class LoadBalancerClientFilter implements GlobalFilter, Ordered { throw new NotFoundException("Unable to find instance for " + url.getHost()); } - /*URI uri = exchange.getRequest().getURI(); - URI requestUrl = loadBalancer.reconstructURI(instance, uri);*/ - boolean encoded = containsEncodedQuery(url); - URI requestUrl = UriComponentsBuilder.fromUri(url) - .scheme(instance.isSecure()? "https" : "http") //TODO: support websockets - .host(instance.getHost()) - .port(instance.getPort()) - .build(encoded) - .toUri(); + URI uri = exchange.getRequest().getURI(); + URI requestUrl = loadBalancer.reconstructURI(instance, uri); + log.trace("LoadBalancerClientFilter url chosen: " + requestUrl); exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, requestUrl); return chain.filter(exchange); diff --git a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilterTests.java b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilterTests.java index 8e8fbeb0..ba80c8a6 100644 --- a/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilterTests.java +++ b/spring-cloud-gateway-core/src/test/java/org/springframework/cloud/gateway/filter/LoadBalancerClientFilterTests.java @@ -1,29 +1,25 @@ -/* - * Copyright 2013-2017 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.Collections; +import java.util.LinkedHashSet; +import com.netflix.loadbalancer.ILoadBalancer; +import com.netflix.loadbalancer.Server; +import org.junit.Before; 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 org.springframework.cloud.client.DefaultServiceInstance; +import org.springframework.cloud.client.ServiceInstance; import org.springframework.cloud.client.loadbalancer.LoadBalancerClient; +import org.springframework.cloud.gateway.support.NotFoundException; +import org.springframework.cloud.netflix.ribbon.RibbonLoadBalancerClient; +import org.springframework.cloud.netflix.ribbon.RibbonLoadBalancerContext; +import org.springframework.cloud.netflix.ribbon.SpringClientFactory; import org.springframework.http.HttpMethod; import org.springframework.mock.http.server.reactive.MockServerHttpRequest; import org.springframework.mock.web.server.MockServerWebExchange; @@ -31,17 +27,103 @@ import org.springframework.web.server.ServerWebExchange; import org.springframework.web.util.UriComponentsBuilder; import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; import static org.mockito.Mockito.mock; +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 reactor.core.publisher.Mono; /** * @author Spencer Gibb + * @author Tim Ysewyn */ +@RunWith(MockitoJUnitRunner.class) public class LoadBalancerClientFilterTests { + private ServerWebExchange exchange; + + @Mock + private GatewayFilterChain chain; + + @Mock + private LoadBalancerClient loadBalancerClient; + + @InjectMocks + private LoadBalancerClientFilter loadBalancerClientFilter; + + @Before + public void setup() { + exchange = MockServerWebExchange.from(MockServerHttpRequest.get("loadbalancerclient.org").build()); + } + + @Test + public void shouldNotFilterWhenGatewayRequestUrlIsMissing() { + loadBalancerClientFilter.filter(exchange, chain); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + verifyZeroInteractions(loadBalancerClient); + } + + @Test + public void shouldNotFilterWhenGatewayRequestUrlSchemeIsNotLb() { + URI uri = UriComponentsBuilder.fromUriString("http://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + + loadBalancerClientFilter.filter(exchange, chain); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + verifyZeroInteractions(loadBalancerClient); + } + + @Test(expected = NotFoundException.class) + public void shouldThrowExceptionWhenNoServiceInstanceIsFound() { + URI uri = UriComponentsBuilder.fromUriString("lb://myservice").build().toUri(); + exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); + + loadBalancerClientFilter.filter(exchange, chain); + } + + @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(loadBalancerClient.choose("myservice")).thenReturn(serviceInstance); + + URI requestUrl = UriComponentsBuilder.fromUriString("https://localhost:8080").build().toUri(); + when(loadBalancerClient.reconstructURI(any(ServiceInstance.class), any(URI.class))).thenReturn(requestUrl); + + loadBalancerClientFilter.filter(exchange, chain); + + assertThat((LinkedHashSet)exchange.getAttribute(GATEWAY_ORIGINAL_REQUEST_URL_ATTR)).contains(url); + + verify(loadBalancerClient).choose("myservice"); + + ArgumentCaptor urlArgumentCaptor = ArgumentCaptor.forClass(URI.class); + verify(loadBalancerClient).reconstructURI(eq(serviceInstance), urlArgumentCaptor.capture()); + + URI uri = urlArgumentCaptor.getValue(); + assertThat(uri).isNotNull(); + assertThat(uri.toString()).isEqualTo("loadbalancerclient.org"); + + verifyNoMoreInteractions(loadBalancerClient); + + assertThat((URI)exchange.getAttribute(GATEWAY_REQUEST_URL_ATTR)).isEqualTo(requestUrl); + + verify(chain).filter(exchange); + verifyNoMoreInteractions(chain); + } + + @Test public void happyPath() { MockServerHttpRequest request = MockServerHttpRequest @@ -119,18 +201,20 @@ public class LoadBalancerClientFilterTests { ServerWebExchange exchange = MockServerWebExchange.from(request); exchange.getAttributes().put(GATEWAY_REQUEST_URL_ATTR, uri); - GatewayFilterChain filterChain = mock(GatewayFilterChain.class); - ArgumentCaptor captor = ArgumentCaptor.forClass(ServerWebExchange.class); - when(filterChain.filter(captor.capture())).thenReturn(Mono.empty()); + when(chain.filter(captor.capture())).thenReturn(Mono.empty()); - LoadBalancerClient loadBalancerClient = mock(LoadBalancerClient.class); - when(loadBalancerClient.choose("service1")). - thenReturn(new DefaultServiceInstance("service1", "service1-host1", 8081, - false, Collections.emptyMap())); - - LoadBalancerClientFilter filter = new LoadBalancerClientFilter(loadBalancerClient); - filter.filter(exchange, filterChain); + SpringClientFactory clientFactory = mock(SpringClientFactory.class); + ILoadBalancer loadBalancer = mock(ILoadBalancer.class); + + when(clientFactory.getLoadBalancerContext("service1")).thenReturn(new RibbonLoadBalancerContext(loadBalancer)); + when(clientFactory.getLoadBalancer("service1")).thenReturn(loadBalancer); + when(loadBalancer.chooseServer(any())).thenReturn(new Server("service1-host1", 8081)); + + RibbonLoadBalancerClient client = new RibbonLoadBalancerClient(clientFactory); + + LoadBalancerClientFilter filter = new LoadBalancerClientFilter(client); + filter.filter(exchange, chain); return captor.getValue(); }