diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagatorGetter.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagatorGetter.java index ac1513f06..ca96e37b0 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagatorGetter.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessageHeaderPropagatorGetter.java @@ -19,6 +19,7 @@ package org.springframework.cloud.sleuth.instrument.messaging; import java.nio.charset.StandardCharsets; import java.util.List; import java.util.Map; +import java.util.Set; import org.apache.commons.logging.Log; import org.apache.commons.logging.LogFactory; @@ -60,26 +61,51 @@ public class MessageHeaderPropagatorGetter implements Propagator.Getter> nativeHeadersMap = nativeAccessor.toNativeHeaderMap(); + if (!nativeHeadersMap.isEmpty()) { + return getFromNativeHeaders(nativeHeadersMap, key); } } else { Object nativeHeaders = accessor.getHeader(NativeMessageHeaderAccessor.NATIVE_HEADERS); if (nativeHeaders instanceof Map) { - Object result = ((Map) nativeHeaders).get(key); - if (result instanceof List && !((List) result).isEmpty()) { - return String.valueOf(((List) result).get(0)); + Map nativeHeadersMap = (Map) nativeHeaders; + if (!nativeHeadersMap.isEmpty()) { + return getFromNativeHeaders(nativeHeadersMap, key); } } } - Object result = accessor.getHeader(key); - if (result != null) { - if (result instanceof byte[]) { - return new String((byte[]) result, StandardCharsets.UTF_8); + Set> headerEntries = accessor.getMessageHeaders().entrySet(); + return getFromHeaders(headerEntries, key); + } + + private String getFromHeaders(Set> headerEntries, String key) { + for (Map.Entry entry : headerEntries) { + if (entry.getKey().equalsIgnoreCase(key)) { + Object result = entry.getValue(); + if (result != null) { + if (result instanceof byte[]) { + return new String((byte[]) result, StandardCharsets.UTF_8); + } + return result.toString(); + } + } + } + return null; + } + + private String getFromNativeHeaders(Map nativeHeaders, String key) { + Set entrySet = nativeHeaders.entrySet(); + for (Map.Entry entries : entrySet) { + if (entries.getKey() instanceof String) { + String headersKey = (String) entries.getKey(); + if (headersKey.equalsIgnoreCase(key)) { + Object result = entries.getValue(); + if (result instanceof List && !((List) result).isEmpty()) { + return String.valueOf(((List) result).get(0)); + } + } } - return result.toString(); } return null; } @@ -88,5 +114,4 @@ public class MessageHeaderPropagatorGetter implements Propagator.Getter !span.equals(initialSpan)) - .allMatch(span -> "FO".equals(COUNTRY_CODE.getValue(BraveAccessor.traceContext(span.context())))); + // it propagates only and all the `spring.sleuth.baggage.remote-fields` in case insensitive way + .allMatch(span -> "FO".equals(COUNTRY_CODE.getValue(BraveAccessor.traceContext(span.context())))) + .allMatch(span -> "123".equalsIgnoreCase(CASE_INSENSITIVE_ID.getValue(BraveAccessor.traceContext(span.context())))) + .allMatch(span -> NOT_PROPAGATED_HEADER.getValue(BraveAccessor.traceContext(span.context())) == null); } @Configuration(proxyBeanMethods = false) diff --git a/tests/brave/spring-cloud-sleuth-instrumentation-messaging-tests/src/test/java/org/springframework/cloud/sleuth/brave/instrument/messaging/TracingChannelInterceptorTest.java b/tests/brave/spring-cloud-sleuth-instrumentation-messaging-tests/src/test/java/org/springframework/cloud/sleuth/brave/instrument/messaging/TracingChannelInterceptorTest.java index 2ff7ef924..7f0ce0993 100644 --- a/tests/brave/spring-cloud-sleuth-instrumentation-messaging-tests/src/test/java/org/springframework/cloud/sleuth/brave/instrument/messaging/TracingChannelInterceptorTest.java +++ b/tests/brave/spring-cloud-sleuth-instrumentation-messaging-tests/src/test/java/org/springframework/cloud/sleuth/brave/instrument/messaging/TracingChannelInterceptorTest.java @@ -20,6 +20,9 @@ import java.util.List; import java.util.Map; import brave.Tracing; +import brave.baggage.BaggageField; +import brave.baggage.BaggagePropagation; +import brave.baggage.BaggagePropagationConfig; import brave.propagation.B3Propagation; import brave.propagation.TraceContext; import org.junit.jupiter.api.Test; @@ -46,7 +49,14 @@ public class TracingChannelInterceptorTest @Override public Tracing.Builder tracingBuilder() { return super.tracingBuilder() - .propagationFactory(B3Propagation.newFactoryBuilder().injectFormat(SINGLE).build()); + .propagationFactory(BaggagePropagation.newFactoryBuilder(B3Propagation.newFactoryBuilder() + .injectFormat(SINGLE) + .build() + ) + .add(BaggagePropagationConfig.SingleBaggageField.remote(BaggageField.create("Foo-Id"))) + .add(BaggagePropagationConfig.SingleBaggageField.remote(BaggageField.create("Baz-Id"))) + .build() + ); } }; this.testTracing.reset(); @@ -67,7 +77,7 @@ public class TracingChannelInterceptorTest TraceContext receiveContext = parseB3SingleFormat( ((List) ((Map) this.channel.receive().getHeaders().get(NATIVE_HEADERS)).get("b3")).get(0).toString()) - .context(); + .context(); assertThat(receiveContext.parentIdString()).isEqualTo("000000000000000b"); } diff --git a/tests/common/src/main/java/org/springframework/cloud/sleuth/baggage/multiple/MultipleHopsIntegrationTests.java b/tests/common/src/main/java/org/springframework/cloud/sleuth/baggage/multiple/MultipleHopsIntegrationTests.java index 3011494ff..83514b534 100644 --- a/tests/common/src/main/java/org/springframework/cloud/sleuth/baggage/multiple/MultipleHopsIntegrationTests.java +++ b/tests/common/src/main/java/org/springframework/cloud/sleuth/baggage/multiple/MultipleHopsIntegrationTests.java @@ -52,7 +52,7 @@ import static org.assertj.core.api.BDDAssertions.then; import static org.awaitility.Awaitility.await; @ContextConfiguration(classes = MultipleHopsIntegrationTests.TestConfig.class) -@TestPropertySource(properties = { "spring.sleuth.baggage.remote-fields=x-vcap-request-id,country-code", +@TestPropertySource(properties = { "spring.sleuth.baggage.remote-fields=x-vcap-request-id,country-code,Foo-Id", "spring.sleuth.baggage.local-fields=bp", "spring.sleuth.integration.enabled=true" }) public abstract class MultipleHopsIntegrationTests { @@ -62,6 +62,10 @@ public abstract class MultipleHopsIntegrationTests { protected static final String COUNTRY_CODE = "country-code"; + protected static final String CASE_INSENSITIVE_ID = "Foo-Id"; + + protected static final String NOT_PROPAGATED_HEADER = "baz-id"; + @Autowired Tracer tracer; @@ -117,6 +121,8 @@ public abstract class MultipleHopsIntegrationTests { // set request ID in a header not with the api explicitly HttpHeaders headers = new HttpHeaders(); headers.put(REQUEST_ID, Collections.singletonList("f4308d05-2228-4468-80f6-92a8377ba193")); + headers.put(CASE_INSENSITIVE_ID, Collections.singletonList("123")); + headers.put(NOT_PROPAGATED_HEADER, Collections.singletonList("456")); RequestEntity requestEntity = new RequestEntity(headers, HttpMethod.GET, URI.create("http://localhost:" + this.testConfig.port + "/greeting")); this.restTemplate.exchange(requestEntity, String.class); diff --git a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java index 323eaff91..d305300b9 100644 --- a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java +++ b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TracingChannelInterceptorTest.java @@ -45,8 +45,10 @@ 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.util.LinkedMultiValueMap; import org.springframework.util.StringUtils; +import static java.util.Collections.singletonList; import static org.assertj.core.api.Assertions.assertThat; import static org.springframework.messaging.support.NativeMessageHeaderAccessor.NATIVE_HEADERS; @@ -348,6 +350,44 @@ public abstract class TracingChannelInterceptorTest implements TestTracingAwareS assertThat(this.spans).extracting(FinishedSpan::getRemoteServiceName).containsOnly("broker", null); } + @Test + public void should_propagate_headers_case_insensitive() { + channel.addInterceptor(this.interceptor); + Map headers = new HashMap<>(); + headers.put("Foo-Id", "123"); + headers.put("baz-id", "456"); + + channel.send(MessageBuilder.createMessage("foo", new MessageHeaders(headers))); + + Message actualMessage = channel.receive(); + + assertThat(actualMessage.getHeaders()).isNotEmpty(); + assertThat(actualMessage.getHeaders().get("not-propagated-header")).isNull(); + assertThat(actualMessage.getHeaders().get("Foo-Id")).isEqualTo("123"); + assertThat(actualMessage.getHeaders().get("baz-id")).isEqualTo("456"); + } + + @Test + public void should_propagate_native_headers_case_insensitive() { + channel.addInterceptor(this.interceptor); + LinkedMultiValueMap nativeHeaders = new LinkedMultiValueMap<>(); + nativeHeaders.put("Foo-Id", singletonList("123")); + nativeHeaders.put("baz-id", singletonList("456")); + Map headers = new HashMap<>(); + headers.put(NATIVE_HEADERS, nativeHeaders); + + channel.send(MessageBuilder.createMessage("foo", new MessageHeaders(headers))); + + Message actualMessage = channel.receive(); + + assertThat(actualMessage.getHeaders()).isNotEmpty(); + LinkedMultiValueMap actualNativeHeaders = (LinkedMultiValueMap) actualMessage.getHeaders().get(NATIVE_HEADERS); + assertThat(actualNativeHeaders).isNotEmpty(); + assertThat(actualNativeHeaders.get("not-propagated-header")).isNull(); + assertThat(actualNativeHeaders.get("Foo-Id")).isEqualTo(singletonList("123")); + assertThat(actualNativeHeaders.get("baz-id")).isEqualTo(singletonList("456")); + } + public ChannelInterceptor producerSideOnly(ChannelInterceptor delegate) { return new ChannelInterceptorAdapter() { @Override