diff --git a/docs/src/main/asciidoc/spring-cloud-commons.adoc b/docs/src/main/asciidoc/spring-cloud-commons.adoc index 0e8e3429..1aaf6cf4 100644 --- a/docs/src/main/asciidoc/spring-cloud-commons.adoc +++ b/docs/src/main/asciidoc/spring-cloud-commons.adoc @@ -968,6 +968,9 @@ public class CustomLoadBalancerConfiguration { } ---- +TIP: For the non-reactive stack, create this supplier with the `withBlockingHealthChecks()`. +You can also pass your own `WebClient` or `RestTemplate` instance to be used for the checks. + WARNING: `HealthCheckServiceInstanceListSupplier` has its own caching mechanism based on Reactor Flux `replay()`. Therefore, if it's being used, you may want to skip wrapping that supplier with `CachingServiceInstanceListSupplier`. === Same instance preference for LoadBalancer diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfiguration.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfiguration.java index 61481e00..44adf557 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfiguration.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfiguration.java @@ -43,6 +43,8 @@ import org.springframework.context.annotation.Primary; import org.springframework.core.annotation.Order; import org.springframework.core.env.Environment; import org.springframework.retry.support.RetryTemplate; +import org.springframework.web.client.RestTemplate; +import org.springframework.web.reactive.function.client.WebClient; /** * @author Spencer Gibb @@ -90,7 +92,7 @@ public class LoadBalancerClientConfiguration { } @Bean - @ConditionalOnBean(ReactiveDiscoveryClient.class) + @ConditionalOnBean({ ReactiveDiscoveryClient.class, WebClient.Builder.class }) @ConditionalOnMissingBean @ConditionalOnProperty(value = "spring.cloud.loadbalancer.configurations", havingValue = "health-check") public ServiceInstanceListSupplier healthCheckDiscoveryClientServiceInstanceListSupplier( @@ -148,12 +150,12 @@ public class LoadBalancerClientConfiguration { } @Bean - @ConditionalOnBean(DiscoveryClient.class) + @ConditionalOnBean({ DiscoveryClient.class, RestTemplate.class }) @ConditionalOnMissingBean @ConditionalOnProperty(value = "spring.cloud.loadbalancer.configurations", havingValue = "health-check") public ServiceInstanceListSupplier healthCheckDiscoveryClientServiceInstanceListSupplier( ConfigurableApplicationContext context) { - return ServiceInstanceListSupplier.builder().withBlockingDiscoveryClient().withHealthChecks() + return ServiceInstanceListSupplier.builder().withBlockingDiscoveryClient().withBlockingHealthChecks() .build(context); } diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplier.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplier.java index 97d3e9a8..6b2cff53 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplier.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplier.java @@ -19,6 +19,7 @@ package org.springframework.cloud.loadbalancer.core; import java.util.ArrayList; import java.util.Collections; import java.util.List; +import java.util.function.BiFunction; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; @@ -31,14 +32,11 @@ import org.springframework.beans.factory.DisposableBean; import org.springframework.beans.factory.InitializingBean; import org.springframework.cloud.client.ServiceInstance; import org.springframework.cloud.client.loadbalancer.LoadBalancerProperties; -import org.springframework.http.HttpStatus; -import org.springframework.web.reactive.function.client.WebClient; -import org.springframework.web.util.UriComponentsBuilder; /** * A {@link ServiceInstanceListSupplier} implementation that verifies whether the - * instances are alive and only returns the healthy one, unless there are none. Uses - * {@link WebClient} to ping the health endpoint of the instances. + * instances are alive and only returns the healthy one, unless there are none. Uses a + * user-provided function to ping the health endpoint of the instances. * * @author Olga Maciaszek-Sharma * @author Roman Matiushchenko @@ -51,19 +49,20 @@ public class HealthCheckServiceInstanceListSupplier extends DelegatingServiceIns private final LoadBalancerProperties.HealthCheck healthCheck; - private final WebClient webClient; - private final String defaultHealthCheckPath; private final Flux> aliveInstancesReplay; private Disposable healthCheckDisposable; + private final BiFunction> aliveFunction; + public HealthCheckServiceInstanceListSupplier(ServiceInstanceListSupplier delegate, - LoadBalancerProperties.HealthCheck healthCheck, WebClient webClient) { + LoadBalancerProperties.HealthCheck healthCheck, + BiFunction> aliveFunction) { super(delegate); defaultHealthCheckPath = healthCheck.getPath().getOrDefault("default", "/actuator/health"); - this.webClient = webClient; + this.aliveFunction = aliveFunction; this.healthCheck = healthCheck; Repeat aliveInstancesReplayRepeat = Repeat .onlyIf(repeatContext -> this.healthCheck.getRefetchInstances()) @@ -129,10 +128,7 @@ public class HealthCheckServiceInstanceListSupplier extends DelegatingServiceIns protected Mono isAlive(ServiceInstance serviceInstance) { String healthCheckPropertyValue = healthCheck.getPath().get(serviceInstance.getServiceId()); String healthCheckPath = healthCheckPropertyValue != null ? healthCheckPropertyValue : defaultHealthCheckPath; - return webClient.get() - .uri(UriComponentsBuilder.fromUri(serviceInstance.getUri()).path(healthCheckPath).build().toUri()) - .exchange().flatMap(clientResponse -> clientResponse.releaseBody() - .thenReturn(HttpStatus.OK.value() == clientResponse.rawStatusCode())); + return aliveFunction.apply(serviceInstance, healthCheckPath); } @Override diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSupplierBuilder.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSupplierBuilder.java index beef94aa..39d20fb2 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSupplierBuilder.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSupplierBuilder.java @@ -16,6 +16,7 @@ package org.springframework.cloud.loadbalancer.core; +import java.net.URI; import java.util.ArrayList; import java.util.List; import java.util.function.BiFunction; @@ -23,6 +24,7 @@ import java.util.function.Function; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import reactor.core.publisher.Mono; import org.springframework.beans.factory.ObjectProvider; import org.springframework.cloud.client.discovery.DiscoveryClient; @@ -31,8 +33,11 @@ import org.springframework.cloud.client.loadbalancer.LoadBalancerProperties; import org.springframework.cloud.loadbalancer.cache.LoadBalancerCacheManager; import org.springframework.cloud.loadbalancer.config.LoadBalancerZoneConfig; import org.springframework.context.ConfigurableApplicationContext; +import org.springframework.http.HttpStatus; import org.springframework.util.Assert; +import org.springframework.web.client.RestTemplate; import org.springframework.web.reactive.function.client.WebClient; +import org.springframework.web.util.UriComponentsBuilder; /** * A Builder for creating a {@link ServiceInstanceListSupplier} hierarchy to be used in @@ -110,7 +115,22 @@ public final class ServiceInstanceListSupplierBuilder { DelegateCreator creator = (context, delegate) -> { LoadBalancerProperties properties = context.getBean(LoadBalancerProperties.class); WebClient.Builder webClient = context.getBean(WebClient.Builder.class); - return new HealthCheckServiceInstanceListSupplier(delegate, properties.getHealthCheck(), webClient.build()); + return healthCheckServiceInstanceListSupplier(webClient.build(), delegate, properties); + }; + this.creators.add(creator); + return this; + } + + /** + * Adds a {@link HealthCheckServiceInstanceListSupplier} that uses user-provided + * {@link WebClient} instance to the {@link ServiceInstanceListSupplier} hierarchy. + * @param webClient a user-provided {@link WebClient} instance + * @return the {@link ServiceInstanceListSupplierBuilder} object + */ + public ServiceInstanceListSupplierBuilder withHealthChecks(WebClient webClient) { + DelegateCreator creator = (context, delegate) -> { + LoadBalancerProperties properties = context.getBean(LoadBalancerProperties.class); + return healthCheckServiceInstanceListSupplier(webClient, delegate, properties); }; this.creators.add(creator); return this; @@ -130,14 +150,29 @@ public final class ServiceInstanceListSupplierBuilder { /** * Adds a {@link HealthCheckServiceInstanceListSupplier} that uses user-provided - * {@link WebClient} instance to the {@link ServiceInstanceListSupplier} hierarchy. - * @param webClient a user-provided {@link WebClient} instance + * {@link RestTemplate} instance to the {@link ServiceInstanceListSupplier} hierarchy. * @return the {@link ServiceInstanceListSupplierBuilder} object */ - public ServiceInstanceListSupplierBuilder withHealthChecks(WebClient webClient) { + public ServiceInstanceListSupplierBuilder withBlockingHealthChecks() { + DelegateCreator creator = (context, delegate) -> { + RestTemplate restTemplate = context.getBean(RestTemplate.class); + LoadBalancerProperties properties = context.getBean(LoadBalancerProperties.class); + return blockingHealthCheckServiceInstanceListSupplier(restTemplate, delegate, properties); + }; + this.creators.add(creator); + return this; + } + + /** + * Adds a {@link HealthCheckServiceInstanceListSupplier} that uses user-provided + * {@link RestTemplate} instance to the {@link ServiceInstanceListSupplier} hierarchy. + * @param restTemplate a user-provided {@link RestTemplate} instance + * @return the {@link ServiceInstanceListSupplierBuilder} object + */ + public ServiceInstanceListSupplierBuilder withBlockingHealthChecks(RestTemplate restTemplate) { DelegateCreator creator = (context, delegate) -> { LoadBalancerProperties properties = context.getBean(LoadBalancerProperties.class); - return new HealthCheckServiceInstanceListSupplier(delegate, properties.getHealthCheck(), webClient); + return blockingHealthCheckServiceInstanceListSupplier(restTemplate, delegate, properties); }; this.creators.add(creator); return this; @@ -225,6 +260,32 @@ public final class ServiceInstanceListSupplierBuilder { return supplier; } + private ServiceInstanceListSupplier healthCheckServiceInstanceListSupplier(WebClient webClient, + ServiceInstanceListSupplier delegate, LoadBalancerProperties properties) { + return new HealthCheckServiceInstanceListSupplier(delegate, properties.getHealthCheck(), + (serviceInstance, healthCheckPath) -> webClient.get() + .uri(UriComponentsBuilder.fromUri(serviceInstance.getUri()).path(healthCheckPath).build() + .toUri()) + .exchange().flatMap(clientResponse -> clientResponse.releaseBody() + .thenReturn(HttpStatus.OK.value() == clientResponse.rawStatusCode()))); + } + + private ServiceInstanceListSupplier blockingHealthCheckServiceInstanceListSupplier(RestTemplate restTemplate, + ServiceInstanceListSupplier delegate, LoadBalancerProperties properties) { + return new HealthCheckServiceInstanceListSupplier(delegate, properties.getHealthCheck(), + (serviceInstance, healthCheckPath) -> Mono.defer(() -> { + URI uri = UriComponentsBuilder.fromUri(serviceInstance.getUri()).path(healthCheckPath).build() + .toUri(); + try { + return Mono + .just(HttpStatus.OK.equals(restTemplate.getForEntity(uri, Void.class).getStatusCode())); + } + catch (Exception ignored) { + return Mono.just(false); + } + })); + } + /** * Allows creating a {@link ServiceInstanceListSupplier} instance based on provided * {@link ConfigurableApplicationContext}. diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfigurationTests.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfigurationTests.java index 6d8becdc..5a356cb5 100644 --- a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfigurationTests.java +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/annotation/LoadBalancerClientConfigurationTests.java @@ -37,6 +37,7 @@ import org.springframework.cloud.loadbalancer.core.ZonePreferenceServiceInstance import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.retry.support.RetryTemplate; +import org.springframework.web.client.RestTemplate; import org.springframework.web.reactive.function.client.WebClient; import static org.assertj.core.api.BDDAssertions.then; @@ -170,6 +171,17 @@ class LoadBalancerClientConfigurationTests { } + @Test + void shouldInstantiateBlockingHealthCheckServiceInstanceListSupplier() { + blockingDiscoveryClientRunner.withUserConfiguration(RestTemplateTestConfig.class) + .withPropertyValues("spring.cloud.loadbalancer.configurations=health-check").run(context -> { + ServiceInstanceListSupplier supplier = context.getBean(ServiceInstanceListSupplier.class); + then(supplier).isInstanceOf(HealthCheckServiceInstanceListSupplier.class); + then(((DelegatingServiceInstanceListSupplier) supplier).getDelegate()) + .isInstanceOf(DiscoveryClientServiceInstanceListSupplier.class); + }); + } + @Configuration protected static class TestConfig { @@ -181,4 +193,14 @@ class LoadBalancerClientConfigurationTests { } + @Configuration + protected static class RestTemplateTestConfig { + + @Bean + RestTemplate restTemplate() { + return new RestTemplate(); + } + + } + } diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/CachingServiceInstanceListSupplierTests.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/CachingServiceInstanceListSupplierTests.java index 726ab021..39ae4107 100644 --- a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/CachingServiceInstanceListSupplierTests.java +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/CachingServiceInstanceListSupplierTests.java @@ -41,6 +41,7 @@ import org.springframework.web.reactive.function.client.WebClient; import static java.time.Duration.ofMillis; import static org.junit.jupiter.api.Assertions.assertTimeoutPreemptively; +import static org.springframework.cloud.loadbalancer.core.ServiceInstanceListSuppliersTestUtils.healthCheckFunction; /** * Tests for {@link CachingServiceInstanceListSupplier}. @@ -140,7 +141,7 @@ class CachingServiceInstanceListSupplierTests { TestHealthCheckServiceInstanceListSupplier(ServiceInstanceListSupplier delegate, LoadBalancerProperties.HealthCheck healthCheck, WebClient webClient) { - super(delegate, healthCheck, webClient); + super(delegate, healthCheck, healthCheckFunction(webClient)); } @Override diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplierTests.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplierTests.java index 0cb59efa..b71d5df2 100644 --- a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplierTests.java +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/HealthCheckServiceInstanceListSupplierTests.java @@ -46,11 +46,13 @@ import org.springframework.context.annotation.Configuration; import org.springframework.test.context.junit.jupiter.SpringExtension; import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.RestController; +import org.springframework.web.client.RestTemplate; import org.springframework.web.reactive.function.client.WebClient; import static org.assertj.core.api.Assertions.assertThat; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; +import static org.springframework.cloud.loadbalancer.core.ServiceInstanceListSuppliersTestUtils.healthCheckFunction; /** * Tests for {@link HealthCheckServiceInstanceListSupplier}. @@ -72,6 +74,8 @@ class HealthCheckServiceInstanceListSupplierTests { private final WebClient webClient = WebClient.create(); + private final RestTemplate restTemplate = new RestTemplate(); + private LoadBalancerProperties.HealthCheck healthCheck; private HealthCheckServiceInstanceListSupplier listSupplier; @@ -97,7 +101,24 @@ class HealthCheckServiceInstanceListSupplierTests { ServiceInstance serviceInstance = new DefaultServiceInstance("ignored-service-1", serviceId, "127.0.0.1", port, false); listSupplier = new HealthCheckServiceInstanceListSupplier( - ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, webClient); + ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, + healthCheckFunction(webClient)); + + boolean alive = listSupplier.isAlive(serviceInstance).block(); + + assertThat(alive).isTrue(); + } + + @SuppressWarnings("ConstantConditions") + @Test + void shouldCheckInstanceWithProvidedHealthCheckPathWithRestTemplate() { + String serviceId = "ignored-service"; + healthCheck.getPath().put("ignored-service", "/health"); + ServiceInstance serviceInstance = new DefaultServiceInstance("ignored-service-1", serviceId, "127.0.0.1", port, + false); + listSupplier = new HealthCheckServiceInstanceListSupplier( + ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, + healthCheckFunction(restTemplate)); boolean alive = listSupplier.isAlive(serviceInstance).block(); @@ -111,7 +132,8 @@ class HealthCheckServiceInstanceListSupplierTests { ServiceInstance serviceInstance = new DefaultServiceInstance("ignored-service-1", serviceId, "127.0.0.1", port, false); listSupplier = new HealthCheckServiceInstanceListSupplier( - ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, webClient); + ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, + healthCheckFunction(webClient)); boolean alive = listSupplier.isAlive(serviceInstance).block(); @@ -126,7 +148,24 @@ class HealthCheckServiceInstanceListSupplierTests { false); healthCheck.getPath().put(serviceId, "/test"); listSupplier = new HealthCheckServiceInstanceListSupplier( - ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, webClient); + ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, + healthCheckFunction(webClient)); + + boolean alive = listSupplier.isAlive(serviceInstance).block(); + + assertThat(alive).isFalse(); + } + + @SuppressWarnings("ConstantConditions") + @Test + void shouldReturnFalseIfEndpointNotFoundWithRestTemplate() { + String serviceId = "ignored-service"; + ServiceInstance serviceInstance = new DefaultServiceInstance("ignored-service-1", serviceId, "127.0.0.1", port, + false); + healthCheck.getPath().put(serviceId, "/test"); + listSupplier = new HealthCheckServiceInstanceListSupplier( + ServiceInstanceListSuppliers.from(serviceId, serviceInstance), healthCheck, + healthCheckFunction(restTemplate)); boolean alive = listSupplier.isAlive(serviceInstance).block(); @@ -151,7 +190,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.doReturn(Mono.just(true)).when(mock).isAlive(serviceInstance1); Mockito.doReturn(Mono.just(false)).when(mock).isAlive(serviceInstance2); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return mock.isAlive(serviceInstance); @@ -180,7 +220,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.doReturn(Mono.just(true)).when(mock).isAlive(serviceInstance1); Mockito.doReturn(Mono.just(true)).when(mock).isAlive(serviceInstance2); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return mock.isAlive(serviceInstance); @@ -210,7 +251,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.doReturn(Mono.just(true)).when(mock).isAlive(serviceInstance1); Mockito.doReturn(Mono.error(new RuntimeException("boom"))).when(mock).isAlive(serviceInstance2); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return mock.isAlive(serviceInstance); @@ -234,7 +276,8 @@ class HealthCheckServiceInstanceListSupplierTests { ServiceInstanceListSupplier delegate = mock(ServiceInstanceListSupplier.class); Mockito.when(delegate.getServiceId()).thenReturn(SERVICE_ID); Mockito.when(delegate.get()).thenReturn(Flux.just(Lists.list(serviceInstance1, serviceInstance2))); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { if (serviceInstance == serviceInstance1) { @@ -261,7 +304,8 @@ class HealthCheckServiceInstanceListSupplierTests { ServiceInstanceListSupplier delegate = mock(ServiceInstanceListSupplier.class); Mockito.when(delegate.getServiceId()).thenReturn(SERVICE_ID); Mockito.when(delegate.get()).thenReturn(Flux.just(Lists.list(serviceInstance1))); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return Mono.just(true); @@ -292,7 +336,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.doReturn(Mono.just(false), Mono.just(true)).when(mock).isAlive(serviceInstance1); Mockito.doReturn(Mono.error(new RuntimeException("boom"))).when(mock).isAlive(serviceInstance2); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return mock.isAlive(serviceInstance); @@ -320,7 +365,8 @@ class HealthCheckServiceInstanceListSupplierTests { HealthCheckServiceInstanceListSupplier mock = mock(HealthCheckServiceInstanceListSupplier.class); Mockito.when(mock.isAlive(serviceInstance1)).thenReturn(Mono.never(), Mono.just(true)); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return mock.isAlive(serviceInstance); @@ -350,7 +396,8 @@ class HealthCheckServiceInstanceListSupplierTests { .delayElements(healthCheck.getInterval().dividedBy(2))); Mockito.when(delegate.get()).thenReturn(instances); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return Mono.just(true); @@ -380,7 +427,8 @@ class HealthCheckServiceInstanceListSupplierTests { ServiceInstanceListSupplier delegate = mock(ServiceInstanceListSupplier.class); when(delegate.get()).thenReturn(Flux.just(Collections.singletonList(serviceInstance1))) .thenReturn(Flux.just(Collections.singletonList(serviceInstance2))); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return Mono.just(true); @@ -405,7 +453,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.when(delegate.getServiceId()).thenReturn(SERVICE_ID); Mockito.when(delegate.get()).thenReturn(Flux.just(Lists.list(serviceInstance1))); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient) { + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)) { @Override protected Mono isAlive(ServiceInstance serviceInstance) { return Mono.just(true); @@ -435,7 +484,8 @@ class HealthCheckServiceInstanceListSupplierTests { Mockito.when(delegate.get()).thenReturn(Flux.>never() .doOnSubscribe(subscription -> subscribed.set(true)).doOnCancel(instancesCanceled::incrementAndGet)); - listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, webClient); + listSupplier = new HealthCheckServiceInstanceListSupplier(delegate, healthCheck, + healthCheckFunction(webClient)); listSupplier.afterPropertiesSet(); diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSuppliersTestUtils.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSuppliersTestUtils.java new file mode 100644 index 00000000..16eff71d --- /dev/null +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/core/ServiceInstanceListSuppliersTestUtils.java @@ -0,0 +1,61 @@ +/* + * Copyright 2013-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.loadbalancer.core; + +import java.net.URI; +import java.util.function.BiFunction; + +import reactor.core.publisher.Mono; + +import org.springframework.cloud.client.ServiceInstance; +import org.springframework.http.HttpStatus; +import org.springframework.web.client.RestTemplate; +import org.springframework.web.reactive.function.client.WebClient; +import org.springframework.web.util.UriComponentsBuilder; + +/** + * A utility class for {@link ServiceInstanceListSupplier} tests. + * + * @author Olga Maciaszek-Sharma + * @since 3.0.0 + */ +final class ServiceInstanceListSuppliersTestUtils { + + private ServiceInstanceListSuppliersTestUtils() { + throw new UnsupportedOperationException("Cannot instantiate utility class"); + } + + static BiFunction> healthCheckFunction(WebClient webClient) { + return (serviceInstance, healthCheckPath) -> webClient.get() + .uri(UriComponentsBuilder.fromUri(serviceInstance.getUri()).path(healthCheckPath).build().toUri()) + .exchange().flatMap(clientResponse -> clientResponse.releaseBody() + .thenReturn(HttpStatus.OK.value() == clientResponse.rawStatusCode())); + } + + static BiFunction> healthCheckFunction(RestTemplate restTemplate) { + return (serviceInstance, healthCheckPath) -> Mono.defer(() -> { + URI uri = UriComponentsBuilder.fromUri(serviceInstance.getUri()).path(healthCheckPath).build().toUri(); + try { + return Mono.just(HttpStatus.OK.equals(restTemplate.getForEntity(uri, Void.class).getStatusCode())); + } + catch (Exception ignored) { + return Mono.just(false); + } + }); + } + +}