diff --git a/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java b/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java index eda0a8cae..ff4990b16 100644 --- a/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java +++ b/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java @@ -110,8 +110,8 @@ class W3CBaggagePropagatorTest { } /** - * We need to use {@link HttpServletRequestWrapper} for the carrier for this test, since it is what combines the - * multiple baggage headers into one. + * We need to use {@link HttpServletRequestWrapper} for the carrier for this test, + * since it is what combines the multiple baggage headers into one. */ @Test void extract_multipleBaggageHeaders() { @@ -120,14 +120,11 @@ class W3CBaggagePropagatorTest { mockRequest.addHeader("baggage", "key2=value2,key3=value3"); HttpServletRequestWrapper carrier = (HttpServletRequestWrapper) HttpServletRequestWrapper.create(mockRequest); - TraceContextOrSamplingFlags contextWithBaggage = propagator - .contextWithBaggage(carrier, context(), HttpServletRequestWrapper::header); + TraceContextOrSamplingFlags contextWithBaggage = propagator.contextWithBaggage(carrier, context(), + HttpServletRequestWrapper::header); Map baggageEntries = BaggageField.getAllValues(contextWithBaggage); - assertThat(baggageEntries) - .hasSize(3) - .containsEntry("key1", "value1") - .containsEntry("key2", "value2") + assertThat(baggageEntries).hasSize(3).containsEntry("key1", "value1").containsEntry("key2", "value2") .containsEntry("key3", "value3"); } diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceHandlerFunction.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceHandlerFunction.java index 8df043c86..4005a505f 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceHandlerFunction.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceHandlerFunction.java @@ -16,7 +16,7 @@ package org.springframework.cloud.sleuth.instrument.web; -import java.util.concurrent.atomic.AtomicReference; +import java.util.Optional; import reactor.core.publisher.Mono; @@ -47,16 +47,15 @@ public class TraceHandlerFunction implements HandlerFunction { @Override public Mono handle(ServerRequest serverRequest) { - AtomicReference scope = new AtomicReference<>(); - return Mono.just(scope) - .doFirst(() -> serverRequest.attribute(TraceWebFilter.TRACE_REQUEST_ATTR) - .ifPresent(span -> scope.set(currentTraceContext().maybeScope(((Span) span).context())))) - .flatMap(r -> this.delegate.handle(serverRequest)).doFinally(signalType -> { - CurrentTraceContext.Scope spanInScope = scope.get(); - if (spanInScope != null) { - spanInScope.close(); - } - }); + Optional spanOptional = serverRequest.attribute(TraceWebFilter.TRACE_REQUEST_ATTR); + if (!spanOptional.isPresent()) { + return this.delegate.handle(serverRequest); + } + return Mono.justOrEmpty(spanOptional).cast(Span.class).flatMap((Span span) -> { + try (CurrentTraceContext.Scope scope = currentTraceContext().maybeScope(span.context())) { + return this.delegate.handle(serverRequest); + } + }); } private CurrentTraceContext currentTraceContext() {