diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptor.java index 408a32481..cc31036a0 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptor.java @@ -22,6 +22,7 @@ import org.springframework.cloud.sleuth.Span; import org.springframework.cloud.sleuth.sampler.NeverSampler; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; +import org.springframework.messaging.MessageDeliveryException; import org.springframework.messaging.MessageHandler; import org.springframework.messaging.support.GenericMessage; import org.springframework.messaging.support.MessageBuilder; @@ -66,7 +67,8 @@ public class TraceChannelInterceptor extends AbstractTraceChannelInterceptor { @Override public Message preSend(Message message, MessageChannel channel) { - MessageBuilder messageBuilder = MessageBuilder.fromMessage(message); + Message retrievedMessage = getMessage(message); + MessageBuilder messageBuilder = MessageBuilder.fromMessage(retrievedMessage); Span parentSpan = getTracer().isTracing() ? getTracer().getCurrentSpan() : buildSpan(new MessagingTextMap(messageBuilder)); String name = getMessageChannelName(channel); @@ -80,7 +82,16 @@ public class TraceChannelInterceptor extends AbstractTraceChannelInterceptor { getSpanInjector().inject(span, new MessagingTextMap(messageBuilder)); MessageHeaderAccessor headers = MessageHeaderAccessor.getMutableAccessor(message); headers.copyHeaders(messageBuilder.build().getHeaders()); - return new GenericMessage(message.getPayload(), headers.getMessageHeaders()); + return new GenericMessage<>(retrievedMessage.getPayload(), headers.getMessageHeaders()); + } + + private Message getMessage(Message message) { + Object payload = message.getPayload(); + if (payload instanceof MessageDeliveryException) { + MessageDeliveryException e = (MessageDeliveryException) payload; + return e.getFailedMessage(); + } + return message; } private Span startSpan(Span span, String name, Message message) { diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java index 5d96da57b..095f23d08 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java @@ -44,6 +44,7 @@ import org.springframework.integration.core.MessagingTemplate; import org.springframework.integration.support.MessageBuilder; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; +import org.springframework.messaging.MessageDeliveryException; import org.springframework.messaging.MessageHandler; import org.springframework.messaging.MessagingException; import org.springframework.messaging.support.ChannelInterceptor; @@ -321,6 +322,25 @@ public class TraceChannelInterceptorTests implements MessageHandler { this.tracedChannel.removeInterceptor(immutableMessageInterceptor); } + @Test + public void workWithMessageDeliveryException() throws Exception { + Message message = new GenericMessage<>(new MessageDeliveryException( + MessageBuilder.withPayload("hi") + .setHeader(TraceMessageHeaders.TRACE_ID_NAME, Span.idToHex(10L)) + .setHeader(TraceMessageHeaders.SPAN_ID_NAME, Span.idToHex(20L)).build() + )); + + this.tracedChannel.send(message); + + String spanId = this.message.getHeaders().get(TraceMessageHeaders.SPAN_ID_NAME, String.class); + then(spanId).isNotNull(); + long traceId = Span + .hexToId(this.message.getHeaders().get(TraceMessageHeaders.TRACE_ID_NAME, String.class)); + then(traceId).isEqualTo(10L); + then(spanId).isNotEqualTo(20L); + then(this.accumulator.getSpans()).hasSize(1); + } + @Configuration @EnableAutoConfiguration static class App {