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 2f847855c..954228760 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 @@ -16,6 +16,9 @@ package org.springframework.cloud.sleuth.instrument.messaging; +import java.util.HashMap; +import java.util.Map; + import org.apache.commons.logging.LogFactory; import org.springframework.beans.factory.BeanFactory; import org.springframework.cloud.sleuth.Log; @@ -27,6 +30,7 @@ import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; import org.springframework.messaging.MessagingException; import org.springframework.messaging.MessageHandler; +import org.springframework.messaging.support.ErrorMessage; import org.springframework.messaging.support.GenericMessage; import org.springframework.messaging.support.MessageBuilder; import org.springframework.messaging.support.MessageHeaderAccessor; @@ -124,10 +128,24 @@ public class TraceChannelInterceptor extends AbstractTraceChannelInterceptor { } getSpanInjector().inject(span, new MessagingTextMap(messageBuilder)); MessageHeaderAccessor headers = MessageHeaderAccessor.getMutableAccessor(message); + if (message instanceof ErrorMessage) { + headers.copyHeaders(sleuthHeaders(messageBuilder.build().getHeaders())); + return new ErrorMessage((Throwable) message.getPayload(), headers.getMessageHeaders()); + } headers.copyHeaders(messageBuilder.build().getHeaders()); return new GenericMessage<>(message.getPayload(), headers.getMessageHeaders()); } + private Map sleuthHeaders(Map headers) { + Map headersToCopy = new HashMap<>(); + for (Map.Entry entry : headers.entrySet()) { + if (TraceMessageHeaders.ALL_HEADERS.contains(entry.getKey())) { + headersToCopy.put(entry.getKey(), entry.getValue()); + } + } + return headersToCopy; + } + private Message getMessage(Message message) { Object payload = message.getPayload(); if (payload instanceof MessagingException) { diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceMessageHeaders.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceMessageHeaders.java index 52af457c0..6767f5ec2 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceMessageHeaders.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceMessageHeaders.java @@ -16,6 +16,9 @@ package org.springframework.cloud.sleuth.instrument.messaging; +import java.util.Arrays; +import java.util.List; + /** * Contains trace related messaging headers. The deprecated headers contained `-` which * for example in the JMS specs is invalid. That's why the public constants in this class @@ -33,6 +36,8 @@ public class TraceMessageHeaders { public static final String TRACE_ID_NAME = "spanTraceId"; public static final String SPAN_NAME_NAME = "spanName"; public static final String SPAN_FLAGS_NAME = "spanFlags"; + static final List ALL_HEADERS = Arrays.asList(SPAN_ID_NAME, SAMPLED_NAME, + PROCESS_ID_NAME, PARENT_ID_NAME, TRACE_ID_NAME, SPAN_NAME_NAME, SPAN_FLAGS_NAME); static final String MESSAGE_SENT_FROM_CLIENT = "messageSent"; static final String HEADER_DELIMITER = "_"; 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 275227bbf..b5676c92d 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 @@ -40,15 +40,17 @@ import org.springframework.cloud.sleuth.util.ExceptionUtils; import org.springframework.context.annotation.Bean; import org.springframework.context.annotation.Configuration; import org.springframework.integration.channel.DirectChannel; +import org.springframework.integration.channel.QueueChannel; 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.MessageHeaders; import org.springframework.messaging.MessagingException; import org.springframework.messaging.support.ChannelInterceptor; import org.springframework.messaging.support.ChannelInterceptorAdapter; +import org.springframework.messaging.support.ErrorMessage; import org.springframework.messaging.support.GenericMessage; import org.springframework.messaging.support.MessageHeaderAccessor; import org.springframework.test.annotation.DirtiesContext; @@ -342,6 +344,35 @@ public class TraceChannelInterceptorTests implements MessageHandler { then(this.accumulator.getSpans()).hasSize(1); } + @Test + public void errorMessageHeadersRetained() { + QueueChannel deadReplyChannel = new QueueChannel(); + QueueChannel errorsReplyChannel = new QueueChannel(); + Map errorChannelHeaders = new HashMap<>(); + errorChannelHeaders.put(MessageHeaders.REPLY_CHANNEL, errorsReplyChannel); + errorChannelHeaders.put(MessageHeaders.ERROR_CHANNEL, errorsReplyChannel); + + this.tracedChannel.send(new ErrorMessage( + new MessagingException(MessageBuilder.withPayload("hi") + .setHeader(TraceMessageHeaders.TRACE_ID_NAME, Span.idToHex(10L)) + .setHeader(TraceMessageHeaders.SPAN_ID_NAME, Span.idToHex(20L)) + .setReplyChannel(deadReplyChannel) + .setErrorChannel(deadReplyChannel) + .build() ), + errorChannelHeaders)); + then(this.message).isNotNull(); + + 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); + then(this.message.getHeaders().getReplyChannel()).isSameAs(errorsReplyChannel); + then(this.message.getHeaders().getErrorChannel()).isSameAs(errorsReplyChannel); + } + @Configuration @EnableAutoConfiguration static class App {