From aea19fcd11075fb8009f9788730886896e02ea2f Mon Sep 17 00:00:00 2001 From: James Moessis Date: Thu, 9 Jun 2022 17:16:22 +1000 Subject: [PATCH] Fixes #2156 (Multiple baggage headers) (#2177) --- .../bridge/W3CBaggagePropagatorTest.java | 22 +++++++++++++++++++ .../servlet/HttpServletRequestWrapper.java | 21 ++++++++++++++++-- 2 files changed, 41 insertions(+), 2 deletions(-) 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..ff4990b16 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,25 @@ 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 133c195ba..2ca125b6b 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; @@ -32,9 +35,15 @@ import org.springframework.lang.Nullable; * @since 5.10 */ // Public for use in sparkjava or other frameworks that re-use servlet types -class HttpServletRequestWrapper implements HttpServerRequest { +public class HttpServletRequestWrapper implements HttpServerRequest { - /** @since 5.10 */ + private static final List COMBINABLE_HEADERS = Collections.singletonList("baggage"); + + /** + * Wraps the request in a tracing representation. + * @param request http request + * @return wrapped request + */ public static HttpServerRequest create(HttpServletRequest request) { return new HttpServletRequestWrapper(request); } @@ -86,6 +95,14 @@ 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); }