diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ReactorSleuth.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ReactorSleuth.java index 4463d8715..5f1c38a42 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ReactorSleuth.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ReactorSleuth.java @@ -19,6 +19,7 @@ package org.springframework.cloud.sleuth.instrument.reactor; import java.util.AbstractQueue; import java.util.Iterator; import java.util.Queue; +import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiConsumer; import java.util.function.BiFunction; import java.util.function.Function; @@ -64,6 +65,8 @@ public abstract class ReactorSleuth { private static final Log log = LogFactory.getLog(ReactorSleuth.class); + private static final String PENDING_SPAN_KEY = "sleuth.pending-span"; + private ReactorSleuth() { } @@ -607,6 +610,47 @@ public abstract class ReactorSleuth { return contextWrappingFunction.apply(context); } + /** + * Retreives the {@link TraceContext} from the current context. + * @param context Reactor context + * @return {@link TraceContext} or {@code null} if none present + */ + @SuppressWarnings("unchecked") + public static TraceContext getParentTraceContext(Context context, TraceContext fallback) { + AtomicReference pendingSpanRef = getPendingSpan(context); + if (pendingSpanRef == null || pendingSpanRef.get() == null) { + return fallback; + } + return pendingSpanRef.get().context(); + } + + /** + * Retreives the pending span from the current context. + * @param context Reactor context + * @return {@code AtomicReference} to span or {@code null} if none present + * @see ReactorSleuth#putPendingSpan(Context, AtomicReference) + */ + @SuppressWarnings("unchecked") + public static AtomicReference getPendingSpan(ContextView context) { + Object objectSpan = context.getOrDefault(ReactorSleuth.PENDING_SPAN_KEY, null); + if ((objectSpan instanceof AtomicReference)) { + return ((AtomicReference) objectSpan); + } + return null; + } + + /** + * Mutates the {@link Context} to include a mutable reference to a span. Can be used + * when you need to mutate the parent operator context with a span created by a child + * operator. + * @param context Reactor context + * @param span atomic reference of a span + * @return mutated context + */ + public static Context putPendingSpan(Context context, AtomicReference span) { + return context.put(PENDING_SPAN_KEY, span); + } + /** * Retrieves span from Reactor context. * @param tracer tracer diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ScopePassingSpanSubscriber.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ScopePassingSpanSubscriber.java index 907788471..14d17abbe 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ScopePassingSpanSubscriber.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/reactor/ScopePassingSpanSubscriber.java @@ -52,12 +52,12 @@ final class ScopePassingSpanSubscriber implements SpanSubscription, Scanna @Nullable TraceContext parent) { this.subscriber = subscriber; this.currentTraceContext = currentTraceContext; - this.parent = parent; - Context context = parent != null && !parent.equals(ctx.getOrDefault(TraceContext.class, null)) - ? ctx.put(TraceContext.class, parent) : ctx; + this.parent = ReactorSleuth.getParentTraceContext(ctx, parent); + Context context = this.parent != null && !this.parent.equals(ctx.getOrDefault(TraceContext.class, null)) + ? ctx.put(TraceContext.class, this.parent) : ctx; this.context = ReactorSleuth.wrapContext(context); if (log.isTraceEnabled()) { - log.trace("Parent span [" + parent + "], context [" + this.context + "]"); + log.trace("Parent span [" + this.parent + "], context [" + this.context + "]"); } } diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessor.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessor.java index 0f2258c4f..ecb766854 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessor.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessor.java @@ -115,7 +115,7 @@ public class HttpClientBeanPostProcessor implements BeanPostProcessor { // Read in this processor and also in ScopePassingSpanSubscriber context = ReactorSleuth.wrapContext(context.put(TraceContext.class, invocationContext)); } - return context.put(PendingSpan.class, pendingSpan); + return ReactorSleuth.putPendingSpan(context, pendingSpan); }).doOnCancel(() -> { // Check to see if Subscription.cancel() happened before another signal, // like onComplete() completed the span (clearing the reference). @@ -153,7 +153,7 @@ public class HttpClientBeanPostProcessor implements BeanPostProcessor { @Override public void accept(HttpClientRequest req, Connection connection) { - PendingSpan pendingSpan = req.currentContextView().getOrDefault(PendingSpan.class, null); + AtomicReference pendingSpan = ReactorSleuth.getPendingSpan(req.currentContextView()); if (pendingSpan == null) { return; // Somehow TracingMapConnect was not invoked.. skip out } @@ -240,7 +240,7 @@ public class HttpClientBeanPostProcessor implements BeanPostProcessor { } void handle(Context context, @Nullable HttpClientResponse resp, @Nullable Throwable error) { - PendingSpan pendingSpan = context.getOrDefault(PendingSpan.class, null); + AtomicReference pendingSpan = ReactorSleuth.getPendingSpan(context); if (pendingSpan == null) { return; // Somehow TracingMapConnect was not invoked.. skip out } diff --git a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessorTest.java b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessorTest.java index ff85ff7ff..e1216d4ab 100644 --- a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessorTest.java +++ b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/web/client/HttpClientBeanPostProcessorTest.java @@ -16,10 +16,13 @@ package org.springframework.cloud.sleuth.instrument.web.client; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicBoolean; import java.util.function.BiConsumer; import io.netty.bootstrap.Bootstrap; import org.assertj.core.api.Assertions; +import org.awaitility.Awaitility; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; @@ -32,7 +35,6 @@ import reactor.core.scheduler.Schedulers; import reactor.netty.Connection; import org.springframework.cloud.sleuth.TraceContext; -import org.springframework.cloud.sleuth.instrument.web.client.HttpClientBeanPostProcessor.PendingSpan; import org.springframework.cloud.sleuth.instrument.web.client.HttpClientBeanPostProcessor.TracingMapConnect; @ExtendWith(MockitoExtension.class) @@ -58,35 +60,51 @@ public abstract class HttpClientBeanPostProcessorTest { @Test void mapConnect_should_setup_reactor_context_currentTraceContext() { TracingMapConnect tracingMapConnect = new TracingMapConnect(() -> traceContext); + AtomicBoolean assertionPassed = new AtomicBoolean(); Mono original = Mono.just(connection) .handle(new BiConsumer>() { @Override public void accept(Connection t, SynchronousSink ctx) { - Assertions.assertThat(ctx.currentContext().get(TraceContext.class)).isSameAs(traceContext); - Assertions.assertThat(ctx.currentContext().get(PendingSpan.class)).isNotNull(); + try { + Assertions.assertThat(ctx.currentContext().get(TraceContext.class)).isSameAs(traceContext); + Assertions.assertThat((Object) ctx.currentContext().get("sleuth.pending-span")).isNotNull(); + assertionPassed.set(true); + } + catch (AssertionError ae) { + } } }); // Wrap and run the assertions tracingMapConnect.apply(original).log().subscribe(); + + Awaitility.await().atMost(1, TimeUnit.SECONDS).untilTrue(assertionPassed); } @Test void mapConnect_should_setup_reactor_context_no_currentTraceContext() { TracingMapConnect tracingMapConnect = new TracingMapConnect(() -> null); + AtomicBoolean assertionPassed = new AtomicBoolean(); Mono original = Mono.just(connection) .handle(new BiConsumer>() { @Override public void accept(Connection t, SynchronousSink ctx) { - Assertions.assertThat(ctx.currentContext().getOrEmpty(TraceContext.class)).isEmpty(); - Assertions.assertThat(ctx.currentContext().get(PendingSpan.class)).isNotNull(); + try { + Assertions.assertThat(ctx.currentContext().getOrEmpty(TraceContext.class)).isEmpty(); + Assertions.assertThat((Object) ctx.currentContext().get("sleuth.pending-span")).isNotNull(); + assertionPassed.set(true); + } + catch (AssertionError ae) { + } } }); // Wrap and run the assertions tracingMapConnect.apply(original).log().subscribe(); + + Awaitility.await().atMost(1, TimeUnit.SECONDS).untilTrue(assertionPassed); } }