diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapper.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapper.java index 63ea201ac..036159d15 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapper.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapper.java @@ -24,6 +24,7 @@ import java.util.concurrent.ConcurrentHashMap; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; +import org.jetbrains.annotations.NotNull; import org.reactivestreams.Publisher; import reactor.core.publisher.Flux; import reactor.core.publisher.Mono; @@ -158,7 +159,16 @@ public class TraceFunctionAroundWrapper extends FunctionAroundWrapper if (targetFunction.isConsumer()) { return targetFunction.apply(reactorStreamConsumer(mono)); } - final Mono function = ((Mono) targetFunction.apply(mono)); + final Publisher function = ((Publisher) targetFunction.apply(mono)); + if (function instanceof Mono) { + return messageMono(targetFunction, (Mono) function); + } + return messageFlux(targetFunction, (Flux) function); + } + + @NotNull + private Mono messageMono(SimpleFunctionRegistry.FunctionInvocationWrapper targetFunction, + Mono function) { return Mono.deferContextual(contextView -> { MessageAndSpansAndScope msg = contextView.get(MessageAndSpansAndScope.class); return function.doOnNext(message -> { @@ -201,7 +211,16 @@ public class TraceFunctionAroundWrapper extends FunctionAroundWrapper if (targetFunction.isConsumer()) { return targetFunction.apply(reactorStreamConsumer(flux)); } - final Flux function = ((Flux) targetFunction.apply(flux)); + final Publisher function = ((Publisher) targetFunction.apply(flux)); + if (function instanceof Mono) { + return messageMono(targetFunction, (Mono) function); + } + return messageFlux(targetFunction, (Flux) function); + } + + @NotNull + private Flux messageFlux(SimpleFunctionRegistry.FunctionInvocationWrapper targetFunction, + Flux function) { return Flux.deferContextual(contextView -> { MessageAndSpansAndScope msg = contextView.get(MessageAndSpansAndScope.class); return function.doOnNext(message -> { diff --git a/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java b/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java index 6fd6ad2fd..3c8a0f17c 100644 --- a/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java +++ b/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java @@ -201,6 +201,25 @@ class TraceFunctionAroundWrapperTests { assertThatAllSpansAreStartedAndStopped(); } + @Test + void should_trace_when_reactive_mono_to_flux_function() { + FunctionRegistration registration = new FunctionRegistration<>( + new ReactiveMonoToFluxFunction(), "greeter").type(FunctionType.of(ReactiveMonoToFluxFunction.class)); + catalog.register(registration); + FunctionInvocationWrapper function = catalog.lookup("greeter"); + + Message result = ((Flux) wrapper.apply( + Mono.just(MessageBuilder.withPayload("hello").setHeader("superHeader", "someValue").build()), function)) + .blockFirst(Duration.ofSeconds(5)); + + assertThat(result.getPayload()).isEqualTo("HELLO"); + assertThat(tracer.spans).hasSize(3); + assertThat(tracer.spans.get(0).name).isEqualTo("handle"); + assertThat(tracer.spans.get(1).name).isEqualTo("greeter"); + assertThat(tracer.spans.get(2).name).isEqualTo("send"); + assertThatAllSpansAreStartedAndStopped(); + } + @Test void should_trace_when_reactive_flux_supplier() { FunctionRegistration registration = new FunctionRegistration<>(new ReactiveFluxGreeter(), @@ -234,6 +253,23 @@ class TraceFunctionAroundWrapperTests { assertThatAllSpansAreStartedAndStopped(); } + @Test + void should_trace_when_reactive_flux_function_returns_mono() { + FunctionRegistration registration = new FunctionRegistration<>( + new ReactiveFluxToMonoFunction(), "greeter").type(FunctionType.of(ReactiveFluxToMonoFunction.class)); + catalog.register(registration); + FunctionInvocationWrapper function = catalog.lookup("greeter"); + + ((Mono) wrapper.apply( + Flux.just(MessageBuilder.withPayload("hello").setHeader("superHeader", "someValue").build()), function)) + .block(Duration.ofSeconds(5)); + + assertThat(tracer.spans).hasSize(2); + assertThat(tracer.spans.get(0).name).isEqualTo("handle"); + assertThat(tracer.spans.get(1).name).isEqualTo("greeter"); + assertThatAllSpansAreStartedAndStopped(); + } + @Test void should_trace_when_reactive_flux_consumer() { ReactiveFluxGreeterConsumer consumer = new ReactiveFluxGreeterConsumer(this.tracer); @@ -352,6 +388,16 @@ class TraceFunctionAroundWrapperTests { } + private static class ReactiveMonoToFluxFunction implements Function>, Flux>> { + + @Override + public Flux> apply(Mono> in) { + return Flux + .from(in.map(s -> MessageBuilder.fromMessage(s).withPayload(s.getPayload().toUpperCase()).build())); + } + + } + private static class ReactiveFluxGreeter implements Supplier>> { @Override @@ -370,6 +416,16 @@ class TraceFunctionAroundWrapperTests { } + private static class ReactiveFluxToMonoFunction implements Function>, Mono> { + + @Override + public Mono apply(Flux> in) { + return in.map(s -> MessageBuilder.fromMessage(s).withPayload(s.getPayload().toUpperCase()).build()) + .then(Mono.empty()); + } + + } + private static class ReactiveFluxGreeterConsumer implements Consumer>> { String result;