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 b6fd97708..342665bed 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 @@ -33,6 +33,7 @@ import org.springframework.integration.context.IntegrationObjectSupport; 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.ChannelInterceptorAdapter; import org.springframework.messaging.support.ErrorMessage; @@ -154,10 +155,17 @@ public final class TracingChannelInterceptor extends ChannelInterceptorAdapter if (originalMessage.getPayload() instanceof MessagingException) { headers.copyHeaders(MessageHeaderPropagation.propagationHeaders(additionalHeaders.getMessageHeaders(), this.tracing.propagation().keys())); - return new ErrorMessage((MessagingException) originalMessage.getPayload(), headers.getMessageHeaders()); + return new ErrorMessage((MessagingException) originalMessage.getPayload(), + isWebSockets(headers) ? headers.getMessageHeaders() : new MessageHeaders(headers.getMessageHeaders())); } headers.copyHeaders(additionalHeaders.getMessageHeaders()); - return new GenericMessage<>(retrievedMessage.getPayload(), headers.getMessageHeaders()); + return new GenericMessage<>(retrievedMessage.getPayload(), + isWebSockets(headers) ? headers.getMessageHeaders() : new MessageHeaders(headers.getMessageHeaders())); + } + + private boolean isWebSockets(MessageHeaderAccessor headerAccessor) { + return headerAccessor.getMessageHeaders().containsKey("stompCommand") || + headerAccessor.getMessageHeaders().containsKey("simpMessageType"); } private boolean isDirectChannel(MessageChannel channel) { diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/ITTracingChannelInterceptor.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/ITTracingChannelInterceptor.java index 3444ed2e0..db36b9b74 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/ITTracingChannelInterceptor.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/ITTracingChannelInterceptor.java @@ -102,13 +102,30 @@ public class ITTracingChannelInterceptor implements MessageHandler { assertThat(currentSpan.isNoop()).isTrue(); } - @Test public void messageHeadersStillMutable() { - directChannel.send(MessageBuilder.withPayload("hi").setHeader("X-B3-Sampled", "0") + @Test public void messageHeadersStillMutableForStomp() { + directChannel.send(MessageBuilder.withPayload("hi").setHeader("stompCommand", "DISCONNECT") .build()); assertThat( MessageHeaderAccessor.getAccessor(message, MessageHeaderAccessor.class)) .isNotNull(); + + message = null; + directChannel.send(MessageBuilder.withPayload("hi").setHeader("simpMessageType", "sth") + .build()); + + assertThat( + MessageHeaderAccessor.getAccessor(message, MessageHeaderAccessor.class)) + .isNotNull(); + } + + @Test public void messageHeadersImmutableForNonStomp() { + directChannel.send(MessageBuilder.withPayload("hi").setHeader("foo", "bar") + .build()); + + assertThat( + MessageHeaderAccessor.getAccessor(message, MessageHeaderAccessor.class)) + .isNull(); } @Configuration @EnableAutoConfiguration static class App {