diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/AbstractTraceChannelInterceptor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/AbstractTraceChannelInterceptor.java index 6f594ea4b..df22ff6e8 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/AbstractTraceChannelInterceptor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/AbstractTraceChannelInterceptor.java @@ -1,5 +1,9 @@ package org.springframework.cloud.sleuth.instrument.messaging; +import java.lang.invoke.MethodHandles; + +import org.apache.commons.logging.Log; +import org.apache.commons.logging.LogFactory; import org.springframework.cloud.sleuth.Span; import org.springframework.cloud.sleuth.SpanExtractor; import org.springframework.cloud.sleuth.SpanInjector; @@ -22,6 +26,8 @@ import org.springframework.util.ClassUtils; abstract class AbstractTraceChannelInterceptor extends ChannelInterceptorAdapter implements ExecutorChannelInterceptor { + private static final Log log = LogFactory.getLog(MethodHandles.lookup().lookupClass()); + /** * If a span comes from messaging components then it will have this value as a prefix * to its name. @@ -63,7 +69,12 @@ abstract class AbstractTraceChannelInterceptor extends ChannelInterceptorAdapter * missing. */ protected Span buildSpan(Message message) { - return this.spanExtractor.joinTrace(message); + try { + return this.spanExtractor.joinTrace(message); + } catch (Exception e) { + log.error("Exception occurred while trying to extract span from carrier", e); + return null; + } } String getChannelName(MessageChannel channel) { diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractor.java index 8081b1104..2a682bfe6 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractor.java @@ -54,13 +54,18 @@ class HttpServletRequestExtractor implements SpanExtractor { // can't build a Span without trace id return null; } - String uri = this.urlPathHelper.getPathWithinApplication(carrier); - boolean skip = this.skipPattern.matcher(uri).matches() - || Span.SPAN_NOT_SAMPLED.equals(carrier.getHeader(Span.SAMPLED_NAME)); - long traceId = Span - .hexToId(carrier.getHeader(Span.TRACE_ID_NAME)); - long spanId = spanId(carrier, traceId); - return buildParentSpan(carrier, uri, skip, traceId, spanId); + try { + String uri = this.urlPathHelper.getPathWithinApplication(carrier); + boolean skip = this.skipPattern.matcher(uri).matches() + || Span.SPAN_NOT_SAMPLED.equals(carrier.getHeader(Span.SAMPLED_NAME)); + long traceId = Span + .hexToId(carrier.getHeader(Span.TRACE_ID_NAME)); + long spanId = spanId(carrier, traceId); + return buildParentSpan(carrier, uri, skip, traceId, spanId); + } catch (Exception e) { + log.error("Exception occurred while trying to extract span from carrier", e); + return null; + } } private long spanId(HttpServletRequest carrier, long traceId) { diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java index 567f90e39..c5a570de8 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/TraceChannelInterceptorTests.java @@ -281,7 +281,7 @@ public class TraceChannelInterceptorTests implements MessageHandler { String traceId = this.message.getHeaders().get(Span.TRACE_ID_NAME, String.class); then(traceId).isNull(); - then(accumulator.getSpans()).isEmpty(); + then(this.accumulator.getSpans()).isEmpty(); then(TestSpanContextHolder.getCurrentSpan()).isNull(); } @@ -299,6 +299,18 @@ public class TraceChannelInterceptorTests implements MessageHandler { then(traceId).isEqualTo(Span.hexToId(lower64Bits)); } + @Test + public void shouldNotBreakWhenInvalidHeadersAreSent() { + this.tracedChannel.send(MessageBuilder.withPayload("hi") + .setHeader(TraceMessageHeaders.PARENT_ID_NAME, "-") + .setHeader(TraceMessageHeaders.TRACE_ID_NAME, Span.idToHex(10L)) + .setHeader(TraceMessageHeaders.SPAN_ID_NAME, Span.idToHex(20L)).build()); + + then(this.message).isNotNull(); + then(this.accumulator.getSpans()).isNotEmpty(); + then(TestSpanContextHolder.getCurrentSpan()).isNull(); + } + @Configuration @EnableAutoConfiguration static class App { diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractorTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractorTests.java index fe1afc5df..6cc85427f 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractorTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/HttpServletRequestExtractorTests.java @@ -16,9 +16,9 @@ package org.springframework.cloud.sleuth.instrument.web; -import javax.servlet.http.HttpServletRequest; import java.util.Random; import java.util.regex.Pattern; +import javax.servlet.http.HttpServletRequest; import org.junit.Before; import org.junit.Test; @@ -28,7 +28,6 @@ import org.mockito.Mock; import org.mockito.runners.MockitoJUnitRunner; import org.springframework.cloud.sleuth.Span; -import static org.assertj.core.api.Assertions.fail; import static org.springframework.cloud.sleuth.assertions.SleuthAssertions.then; @RunWith(MockitoJUnitRunner.class) @@ -54,12 +53,7 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.TRACE_ID_NAME)) .willReturn("invalid"); - try { - this.extractor.joinTrace(this.request); - fail("should throw an exception"); - } catch (IllegalArgumentException e) { - then(e).hasMessageContaining("Malformed id"); - } + then(this.extractor.joinTrace(this.request)).isNull(); } @Test @@ -69,12 +63,7 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.SPAN_ID_NAME)) .willReturn("invalid"); - try { - this.extractor.joinTrace(this.request); - fail("should throw an exception"); - } catch (IllegalArgumentException e) { - then(e).hasMessageContaining("Malformed id"); - } + then(this.extractor.joinTrace(this.request)).isNull(); } @Test @@ -86,12 +75,7 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.PARENT_ID_NAME)) .willReturn("invalid"); - try { - this.extractor.joinTrace(this.request); - fail("should throw an exception"); - } catch (IllegalArgumentException e) { - then(e).hasMessageContaining("Malformed id"); - } + then(this.extractor.joinTrace(this.request)).isNull(); } @Test diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterTests.java index 4787beef3..a06f90198 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterTests.java @@ -16,10 +16,11 @@ package org.springframework.cloud.sleuth.instrument.web; -import javax.servlet.http.HttpServletRequest; +import java.util.ArrayList; import java.util.Optional; import java.util.Random; import java.util.regex.Pattern; +import javax.servlet.http.HttpServletRequest; import org.junit.After; import org.junit.Before; @@ -38,6 +39,7 @@ import org.springframework.cloud.sleuth.sampler.NeverSampler; import org.springframework.cloud.sleuth.trace.DefaultTracer; import org.springframework.cloud.sleuth.trace.TestSpanContextHolder; import org.springframework.cloud.sleuth.util.ArrayListSpanAccumulator; +import org.springframework.cloud.sleuth.util.ExceptionUtils; import org.springframework.http.HttpStatus; import org.springframework.http.MediaType; import org.springframework.mock.web.MockFilterChain; @@ -323,7 +325,7 @@ public class TraceFilterTests { } @Test - public void returns400IfSpanIsMalformed() throws Exception { + public void returns400IfSpanIsMalformedAndCreatesANewSpan() throws Exception { this.request = builder().header(Span.SPAN_ID_NAME, "asd") .header(Span.TRACE_ID_NAME, 20L).buildRequest(new MockServletContext()); TraceFilter filter = new TraceFilter(this.tracer, this.traceKeys, this.spanReporter, @@ -331,8 +333,26 @@ public class TraceFilterTests { filter.doFilter(this.request, this.response, this.filterChain); + then(new ArrayList<>(this.spanReporter.getSpans())).isNotEmpty(); then(TestSpanContextHolder.getCurrentSpan()).isNull(); - then(this.response.getStatus()).isEqualTo(HttpStatus.BAD_REQUEST.value()); + then(ExceptionUtils.getLastException()).isNull(); + then(this.response.getStatus()).isEqualTo(HttpStatus.OK.value()); + } + + @Test + public void returns200IfSpanParentIsMalformedAndCreatesANewSpan() throws Exception { + this.request = builder().header(Span.SPAN_ID_NAME, PARENT_ID) + .header(Span.PARENT_ID_NAME, "-") + .header(Span.TRACE_ID_NAME, 20L).buildRequest(new MockServletContext()); + TraceFilter filter = new TraceFilter(this.tracer, this.traceKeys, this.spanReporter, + this.spanExtractor, this.httpTraceKeysInjector); + + filter.doFilter(this.request, this.response, this.filterChain); + + then(new ArrayList<>(this.spanReporter.getSpans())).isNotEmpty(); + then(TestSpanContextHolder.getCurrentSpan()).isNull(); + then(ExceptionUtils.getLastException()).isNull(); + then(this.response.getStatus()).isEqualTo(HttpStatus.OK.value()); } public void verifyParentSpanHttpTags() {