From ae1ac8b005a57d2fbb5e90cb3b340e24b1ca15a9 Mon Sep 17 00:00:00 2001 From: Olga Maciaszek-Sharma Date: Wed, 22 Jan 2025 18:41:55 +0100 Subject: [PATCH] Handle null in LoadBalancer Request and Response (#1454) * Handle null lbRequest and lbResponse values. Signed-off-by: Olga Maciaszek-Sharma * Verify for null request on start. Signed-off-by: Olga Maciaszek-Sharma --------- Signed-off-by: Olga Maciaszek-Sharma --- .../loadbalancer/stats/LoadBalancerTags.java | 26 +++++++++++++------ .../MicrometerStatsLoadBalancerLifecycle.java | 16 +++++++++--- ...ometerStatsLoadBalancerLifecycleTests.java | 21 +++++++++++++++ 3 files changed, 51 insertions(+), 12 deletions(-) diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/LoadBalancerTags.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/LoadBalancerTags.java index 42afc64f..9683fb73 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/LoadBalancerTags.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/LoadBalancerTags.java @@ -27,8 +27,10 @@ import io.micrometer.core.instrument.Tags; import org.springframework.cloud.client.ServiceInstance; import org.springframework.cloud.client.loadbalancer.CompletionContext; import org.springframework.cloud.client.loadbalancer.LoadBalancerProperties; +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.util.StringUtils; @@ -55,7 +57,11 @@ class LoadBalancerTags { } Iterable buildSuccessRequestTags(CompletionContext completionContext) { - ServiceInstance serviceInstance = completionContext.getLoadBalancerResponse().getServer(); + Response lbResponse = completionContext.getLoadBalancerResponse(); + if (lbResponse == null) { + return Tags.empty(); + } + ServiceInstance serviceInstance = lbResponse.getServer(); Tags tags = Tags.of(buildServiceInstanceTags(serviceInstance)); Object clientResponse = completionContext.getClientResponse(); if (clientResponse instanceof ResponseData responseData) { @@ -100,9 +106,9 @@ class LoadBalancerTags { } Iterable buildDiscardedRequestTags(CompletionContext completionContext) { - if (completionContext.getLoadBalancerRequest().getContext() instanceof RequestDataContext) { - RequestData requestData = ((RequestDataContext) completionContext.getLoadBalancerRequest().getContext()) - .getClientRequest(); + Request lbRequest = completionContext.getLoadBalancerRequest(); + if (lbRequest != null && lbRequest.getContext() instanceof RequestDataContext requestDataContext) { + RequestData requestData = requestDataContext.getClientRequest(); if (requestData != null) { return Tags.of(valueOrUnknown("method", requestData.getHttpMethod()), valueOrUnknown("uri", getPath(requestData)), valueOrUnknown("serviceId", getHost(requestData))); @@ -118,11 +124,15 @@ class LoadBalancerTags { } Iterable buildFailedRequestTags(CompletionContext completionContext) { - ServiceInstance serviceInstance = completionContext.getLoadBalancerResponse().getServer(); + Response lbResponse = completionContext.getLoadBalancerResponse(); + if (lbResponse == null) { + return Tags.empty(); + } + ServiceInstance serviceInstance = lbResponse.getServer(); Tags tags = Tags.of(buildServiceInstanceTags(serviceInstance)).and(exception(completionContext.getThrowable())); - if (completionContext.getLoadBalancerRequest().getContext() instanceof RequestDataContext) { - RequestData requestData = ((RequestDataContext) completionContext.getLoadBalancerRequest().getContext()) - .getClientRequest(); + Request lbRequest = completionContext.getLoadBalancerRequest(); + if (lbRequest != null && lbRequest.getContext() instanceof RequestDataContext requestDataContext) { + RequestData requestData = requestDataContext.getClientRequest(); if (requestData != null) { return tags.and(Tags.of(valueOrUnknown("method", requestData.getHttpMethod()), valueOrUnknown("uri", getPath(requestData)))); diff --git a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycle.java b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycle.java index 1d77caea..7b48e9e5 100644 --- a/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycle.java +++ b/spring-cloud-loadbalancer/src/main/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycle.java @@ -85,10 +85,10 @@ public class MicrometerStatsLoadBalancerLifecycle implements LoadBalancerLifecyc @Override public void onStartRequest(Request request, Response lbResponse) { - if (request.getContext() instanceof TimedRequestContext) { + if (request != null && request.getContext() instanceof TimedRequestContext) { ((TimedRequestContext) request.getContext()).setRequestStartTime(System.nanoTime()); } - if (!lbResponse.hasServer()) { + if (lbResponse == null || !lbResponse.hasServer()) { return; } ServiceInstance serviceInstance = lbResponse.getServer(); @@ -104,7 +104,11 @@ public class MicrometerStatsLoadBalancerLifecycle implements LoadBalancerLifecyc @Override public void onComplete(CompletionContext completionContext) { - ServiceInstance serviceInstance = completionContext.getLoadBalancerResponse().getServer(); + ServiceInstance serviceInstance = null; + Response loadBalancerResponse = completionContext.getLoadBalancerResponse(); + if (loadBalancerResponse != null) { + serviceInstance = loadBalancerResponse.getServer(); + } LoadBalancerProperties properties = serviceInstance != null ? loadBalancerFactory.getProperties(serviceInstance.getServiceId()) : loadBalancerFactory.getProperties(null); @@ -121,7 +125,11 @@ public class MicrometerStatsLoadBalancerLifecycle implements LoadBalancerLifecyc if (activeRequestsCounter != null) { activeRequestsCounter.decrementAndGet(); } - Object loadBalancerRequestContext = completionContext.getLoadBalancerRequest().getContext(); + Request lbRequest = completionContext.getLoadBalancerRequest(); + if (lbRequest == null) { + return; + } + Object loadBalancerRequestContext = lbRequest.getContext(); if (requestHasBeenTimed(loadBalancerRequestContext)) { if (CompletionContext.Status.FAILED.equals(completionContext.status())) { Timer.builder("loadbalancer.requests.failed") diff --git a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycleTests.java b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycleTests.java index 27d360dd..3be772cc 100644 --- a/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycleTests.java +++ b/spring-cloud-loadbalancer/src/test/java/org/springframework/cloud/loadbalancer/stats/MicrometerStatsLoadBalancerLifecycleTests.java @@ -47,6 +47,7 @@ import org.springframework.http.HttpStatus; import org.springframework.util.MultiValueMapAdapter; import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatCode; import static org.mockito.Mockito.mock; import static org.mockito.Mockito.when; import static org.springframework.cloud.loadbalancer.stats.LoadBalancerTags.UNKNOWN; @@ -243,6 +244,26 @@ class MicrometerStatsLoadBalancerLifecycleTests { Tag.of("serviceInstance.port", "0"), Tag.of("status", "200"), Tag.of("uri", UNKNOWN)); } + @Test + void shouldHandleNullLoadBalancerResponse() { + RequestData requestData = new RequestData(HttpMethod.GET, URI.create("http://test.org/test"), new HttpHeaders(), + new HttpHeaders(), new HashMap<>()); + Request lbRequest = new DefaultRequest<>(new RequestDataContext(requestData)); + assertThatCode(() -> { + statsLifecycle.onStartRequest(lbRequest, null); + statsLifecycle.onComplete(new CompletionContext<>(CompletionContext.Status.DISCARD, lbRequest, null)); + }).doesNotThrowAnyException(); + } + + @Test + void shouldHandleNullLoadBalancerRequest() { + Response lbResponse = new EmptyResponse(); + assertThatCode(() -> { + statsLifecycle.onStartRequest(null, lbResponse); + statsLifecycle.onComplete(new CompletionContext<>(CompletionContext.Status.DISCARD, null, lbResponse)); + }).doesNotThrowAnyException(); + } + private static class StatsTestContext { }