diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagation.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagation.java index 47aedf2fa..3feedf52e 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagation.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagation.java @@ -137,6 +137,17 @@ enum MessageHeaderPropagation return null; } + static Map propagationHeaders(Map headers, + List propagationHeaders) { + Map headersToCopy = new HashMap<>(); + for (Map.Entry entry : headers.entrySet()) { + if (propagationHeaders.contains(entry.getKey())) { + headersToCopy.put(entry.getKey(), entry.getValue()); + } + } + return headersToCopy; + } + static void removeAnyTraceHeaders(MessageHeaderAccessor accessor, List keysToRemove) { for (String keyToRemove : keysToRemove) { diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptor.java index f44bceddf..8f45fe7f9 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptor.java @@ -120,7 +120,8 @@ public final class TracingChannelInterceptor extends ChannelInterceptorAdapter if (emptyMessage(message)) { return message; } - MessageHeaderAccessor headers = mutableHeaderAccessor(message); + Message retrievedMessage = getMessage(message); + MessageHeaderAccessor headers = mutableHeaderAccessor(retrievedMessage); TraceContextOrSamplingFlags extracted = this.extractor.extract(headers); Span span = this.threadLocalSpan.next(extracted); MessageHeaderPropagation @@ -134,14 +135,24 @@ public final class TracingChannelInterceptor extends ChannelInterceptorAdapter if (log.isDebugEnabled()) { log.debug("Created a new span in pre send" + span); } - headers.setImmutable(); - Message outputMessage = new GenericMessage<>(message.getPayload(), headers.getMessageHeaders()); + Message outputMessage = outputMessage(message, retrievedMessage, headers); if (isDirectChannel(channel)) { beforeHandle(outputMessage, channel, null); } return outputMessage; } + private Message outputMessage(Message originalMessage, Message retrievedMessage, MessageHeaderAccessor additionalHeaders) { + MessageHeaderAccessor headers = MessageHeaderAccessor.getMutableAccessor(originalMessage); + if (originalMessage.getPayload() instanceof MessagingException) { + headers.copyHeaders(MessageHeaderPropagation.propagationHeaders(additionalHeaders.getMessageHeaders(), + this.tracing.propagation().keys())); + return new ErrorMessage((MessagingException) originalMessage.getPayload(), headers.getMessageHeaders()); + } + headers.copyHeaders(additionalHeaders.getMessageHeaders()); + return new GenericMessage<>(retrievedMessage.getPayload(), headers.getMessageHeaders()); + } + private boolean isDirectChannel(MessageChannel channel) { return DirectChannel.class .isAssignableFrom(AopUtils.getTargetClass(channel)); @@ -292,7 +303,7 @@ public final class TracingChannelInterceptor extends ChannelInterceptorAdapter } private MessageHeaderAccessor mutableHeaderAccessor(Message message) { - MessageHeaderAccessor headers = MessageHeaderAccessor.getMutableAccessor(getMessage(message)); + MessageHeaderAccessor headers = MessageHeaderAccessor.getMutableAccessor(message); headers.setLeaveMutable(true); return headers; } diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java index df29739e6..24b5bc48d 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java @@ -18,26 +18,29 @@ package org.springframework.cloud.sleuth.instrument.messaging; import java.util.ArrayList; import java.util.Collections; +import java.util.HashMap; import java.util.List; import java.util.Map; import brave.Tracing; import brave.propagation.StrictCurrentTraceContext; -import org.springframework.integration.channel.DirectChannel; -import org.springframework.messaging.MessagingException; -import zipkin2.Span; import org.junit.After; import org.junit.Test; +import org.springframework.integration.channel.DirectChannel; import org.springframework.integration.channel.QueueChannel; import org.springframework.messaging.Message; import org.springframework.messaging.MessageChannel; 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.ExecutorChannelInterceptor; import org.springframework.messaging.support.ExecutorSubscribableChannel; import org.springframework.messaging.support.MessageBuilder; import org.springframework.messaging.support.NativeMessageHeaderAccessor; +import zipkin2.Span; import static org.assertj.core.api.Assertions.assertThat; import static org.springframework.messaging.support.NativeMessageHeaderAccessor.NATIVE_HEADERS; @@ -226,6 +229,36 @@ public class TracingChannelInterceptorTest { .containsExactly(Span.Kind.CONSUMER, null, Span.Kind.PRODUCER); } + @Test + public void errorMessageHeadersRetained() { + this.channel.addInterceptor(interceptor); + 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.channel.send(new ErrorMessage( + new MessagingException(MessageBuilder.withPayload("hi") + .setHeader(TraceMessageHeaders.TRACE_ID_NAME, "000000000000000a") + .setHeader(TraceMessageHeaders.SPAN_ID_NAME, "000000000000000a") + .setReplyChannel(deadReplyChannel) + .setErrorChannel(deadReplyChannel) + .build()), + errorChannelHeaders)); + + this.message = this.channel.receive(); + + assertThat(this.message).isNotNull(); + String spanId = this.message.getHeaders().get(TraceMessageHeaders.SPAN_ID_NAME, String.class); + assertThat(spanId).isNotNull(); + String traceId = this.message.getHeaders().get(TraceMessageHeaders.TRACE_ID_NAME, String.class); + assertThat(traceId).isEqualTo("000000000000000a"); + assertThat(spanId).isNotEqualTo("000000000000000a"); + assertThat(this.spans).hasSize(2); + assertThat(this.message.getHeaders().getReplyChannel()).isSameAs(errorsReplyChannel); + assertThat(this.message.getHeaders().getErrorChannel()).isSameAs(errorsReplyChannel); + } + ChannelInterceptor producerSideOnly(ChannelInterceptor delegate) { return new ChannelInterceptorAdapter() { @Override