diff --git a/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java b/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java index b997013c7..eda0a8cae 100644 --- a/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java +++ b/spring-cloud-sleuth-brave/src/test/java/org/springframework/cloud/sleuth/brave/bridge/W3CBaggagePropagatorTest.java @@ -29,6 +29,9 @@ import brave.propagation.TraceContextOrSamplingFlags; import org.junit.jupiter.api.Disabled; import org.junit.jupiter.api.Test; +import org.springframework.cloud.sleuth.instrument.web.servlet.HttpServletRequestWrapper; +import org.springframework.mock.web.MockHttpServletRequest; + import static java.util.Collections.singletonMap; import static org.assertj.core.api.Assertions.assertThat; @@ -106,6 +109,28 @@ class W3CBaggagePropagatorTest { assertThat(baggageEntries).hasSize(1).containsEntry("key", "value2"); } + /** + * We need to use {@link HttpServletRequestWrapper} for the carrier for this test, since it is what combines the + * multiple baggage headers into one. + */ + @Test + void extract_multipleBaggageHeaders() { + MockHttpServletRequest mockRequest = new MockHttpServletRequest(); + mockRequest.addHeader("baggage", "key1=value1;othermetadata"); + mockRequest.addHeader("baggage", "key2=value2,key3=value3"); + HttpServletRequestWrapper carrier = (HttpServletRequestWrapper) HttpServletRequestWrapper.create(mockRequest); + + TraceContextOrSamplingFlags contextWithBaggage = propagator + .contextWithBaggage(carrier, context(), HttpServletRequestWrapper::header); + + Map baggageEntries = BaggageField.getAllValues(contextWithBaggage); + assertThat(baggageEntries) + .hasSize(3) + .containsEntry("key1", "value1") + .containsEntry("key2", "value2") + .containsEntry("key3", "value3"); + } + @Test void extract_fullComplexities() { TraceContextOrSamplingFlags context = context(); diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/servlet/HttpServletRequestWrapper.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/servlet/HttpServletRequestWrapper.java index a692c6d5f..6abb376f6 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/servlet/HttpServletRequestWrapper.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/servlet/HttpServletRequestWrapper.java @@ -18,6 +18,9 @@ package org.springframework.cloud.sleuth.instrument.web.servlet; import java.util.Collection; import java.util.Collections; +import java.util.Enumeration; +import java.util.LinkedList; +import java.util.List; import javax.servlet.RequestDispatcher; import javax.servlet.http.HttpServletRequest; @@ -33,6 +36,8 @@ import org.springframework.lang.Nullable; */ public class HttpServletRequestWrapper implements HttpServerRequest { + private static final List COMBINABLE_HEADERS = Collections.singletonList("baggage"); + /** * Wraps the request in a tracing representation. * @param request http request @@ -89,6 +94,14 @@ public class HttpServletRequestWrapper implements HttpServerRequest { @Override public String header(String name) { + if (COMBINABLE_HEADERS.contains(name)) { + LinkedList headersList = new LinkedList<>(); + Enumeration headers = delegate.getHeaders(name); + while (headers.hasMoreElements()) { + headersList.add(headers.nextElement()); + } + return headersList.size() != 0 ? String.join(",", headersList) : null; + } return delegate.getHeader(name); }