diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/Span.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/Span.java index 4fe24b6e4..5034aad0f 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/Span.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/Span.java @@ -371,7 +371,11 @@ public class Span { */ public static long hexToId(String hexString) { Assert.hasText(hexString, "Can't convert empty hex string to long"); - return new BigInteger(hexString, 16).longValue(); + try { + return new BigInteger(hexString, 16).longValue(); + } catch (NumberFormatException e) { + throw new IllegalArgumentException("Malformed id [" + hexString + "]", e); + } } @Override diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractor.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractor.java index a51c9cf7b..92bd1a780 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractor.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractor.java @@ -16,11 +16,8 @@ package org.springframework.cloud.sleuth.instrument.messaging; -import java.lang.invoke.MethodHandles; import java.util.Random; -import org.apache.commons.logging.Log; -import org.apache.commons.logging.LogFactory; import org.springframework.cloud.sleuth.Span; import org.springframework.cloud.sleuth.Span.SpanBuilder; import org.springframework.cloud.sleuth.SpanExtractor; @@ -34,8 +31,6 @@ import org.springframework.messaging.Message; */ public class MessagingSpanExtractor implements SpanExtractor> { - private static final Log log = LogFactory.getLog(MethodHandles.lookup().lookupClass()); - private final Random random; public MessagingSpanExtractor(Random random) { @@ -49,9 +44,10 @@ public class MessagingSpanExtractor implements SpanExtractor> { return null; // TODO: Consider throwing IllegalArgumentException; } - long traceId = getTraceIdOrSetDefault(carrier); + long traceId = Span + .hexToId(getHeader(carrier, Span.TRACE_ID_NAME)); long spanId = hasHeader(carrier, Span.SPAN_ID_NAME) - ? getSpanIdOrSetDefault(carrier) + ? Span.hexToId(getHeader(carrier, Span.SPAN_ID_NAME)) : this.random.nextLong(); SpanBuilder spanBuilder = Span.builder().traceId(traceId).spanId(spanId); spanBuilder.exportable( @@ -81,39 +77,10 @@ public class MessagingSpanExtractor implements SpanExtractor> { return message.getHeaders().containsKey(name); } - private long getTraceIdOrSetDefault(Message carrier) { - try { - return Span - .hexToId(getHeader(carrier, Span.TRACE_ID_NAME)); - } catch (Exception e) { - long id = this.random.nextLong(); - log.warn("Exception occurred while trying to retrieve the trace " - + "id from headers. Will set id to value [" - + Span.idToHex(id) + "]", e); - return id; - } - } private void setParentIdIfApplicable(Message carrier, SpanBuilder spanBuilder) { - try { - String parentId = getHeader(carrier, Span.PARENT_ID_NAME); - if (parentId != null) { - spanBuilder.parent(Span.hexToId(parentId)); - } - } catch (Exception e) { - log.warn("Exception occurred while trying to set parentId", e); - } - } - - private long getSpanIdOrSetDefault(Message carrier) { - try { - return Span - .hexToId(getHeader(carrier, Span.SPAN_ID_NAME)); - } catch (Exception e) { - long id = this.random.nextLong(); - log.warn("Exception occurred while trying to retrieve the span " - + "id from headers. Will set id to value [" - + Span.idToHex(id) + "]", e); - return id; + String parentId = getHeader(carrier, Span.PARENT_ID_NAME); + if (parentId != null) { + spanBuilder.parent(Span.hexToId(parentId)); } } } 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 45527099b..d0d0f9f21 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 @@ -18,7 +18,6 @@ package org.springframework.cloud.sleuth.instrument.web; import javax.servlet.http.HttpServletRequest; import java.lang.invoke.MethodHandles; -import java.util.Random; import java.util.regex.Pattern; import org.apache.commons.logging.Log; @@ -42,13 +41,11 @@ class HttpServletRequestExtractor implements SpanExtractor { private static final String HTTP_COMPONENT = "http"; private final Pattern skipPattern; - private final Random random; private UrlPathHelper urlPathHelper = new UrlPathHelper(); - public HttpServletRequestExtractor(Pattern skipPattern, Random random) { + public HttpServletRequestExtractor(Pattern skipPattern) { this.skipPattern = skipPattern; - this.random = random; } @Override @@ -60,24 +57,12 @@ class HttpServletRequestExtractor implements SpanExtractor { 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 = getTraceIdOrSetDefault(carrier); + long traceId = Span + .hexToId(carrier.getHeader(Span.TRACE_ID_NAME)); long spanId = spanId(carrier, traceId); return buildParentSpan(carrier, uri, skip, traceId, spanId); } - private long getTraceIdOrSetDefault(HttpServletRequest carrier) { - try { - return Span - .hexToId(carrier.getHeader(Span.TRACE_ID_NAME)); - } catch (Exception e) { - long id = this.random.nextLong(); - log.warn("Exception occurred while trying to retrieve the trace " - + "id from headers. Will set id to value [" - + Span.idToHex(id) + "]", e); - return id; - } - } - private long spanId(HttpServletRequest carrier, long traceId) { String spanId = carrier.getHeader(Span.SPAN_ID_NAME); if (spanId == null) { @@ -85,15 +70,7 @@ class HttpServletRequestExtractor implements SpanExtractor { + "a root span with span id equal to trace id"); return traceId; } else { - try { - return Span.hexToId(spanId); - } catch (Exception e) { - long id = this.random.nextLong(); - log.warn("Exception occurred while trying to retrieve the span id " - + "from request headers. Will set id to value [" - + Span.idToHex(id) + "]", e); - return id; - } + return Span.hexToId(spanId); } } @@ -112,7 +89,8 @@ class HttpServletRequestExtractor implements SpanExtractor { span.processId(processId); } if (carrier.getHeader(Span.PARENT_ID_NAME) != null) { - setParentIdIfValid(carrier, span); + span.parent(Span + .hexToId(carrier.getHeader(Span.PARENT_ID_NAME))); } span.remote(true); if (skip) { @@ -120,13 +98,4 @@ class HttpServletRequestExtractor implements SpanExtractor { } return span.build(); } - - private void setParentIdIfValid(HttpServletRequest carrier, SpanBuilder span) { - try { - span.parent(Span - .hexToId(carrier.getHeader(Span.PARENT_ID_NAME))); - } catch (Exception e) { - log.warn("Exception occurred while trying to set parent id", e); - } - } } diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceFilter.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceFilter.java index cc1db1df8..d5884fc61 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceFilter.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceFilter.java @@ -152,7 +152,14 @@ public class TraceFilter extends GenericFilterBean { } addToResponseIfNotPresent(response, Span.SAMPLED_NAME, skip ? Span.SPAN_NOT_SAMPLED : Span.SPAN_SAMPLED); String name = HTTP_COMPONENT + ":" + uri; - spanFromRequest = createSpan(request, skip, spanFromRequest, name); + try { + spanFromRequest = createSpan(request, skip, spanFromRequest, name); + } catch (IllegalArgumentException e) { + filterChain.doFilter(request, response); + response.sendError(HttpStatus.BAD_REQUEST.value(), + "Exception tracing request [" + e.getMessage() + "]"); + return; + } Throwable exception = null; try { this.spanInjector.inject(spanFromRequest, response); diff --git a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceWebAutoConfiguration.java b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceWebAutoConfiguration.java index 763850d63..47c5bee0d 100644 --- a/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceWebAutoConfiguration.java +++ b/spring-cloud-sleuth-core/src/main/java/org/springframework/cloud/sleuth/instrument/web/TraceWebAutoConfiguration.java @@ -15,7 +15,6 @@ */ package org.springframework.cloud.sleuth.instrument.web; -import java.util.Random; import java.util.regex.Pattern; import javax.servlet.http.HttpServletRequest; @@ -101,8 +100,8 @@ public class TraceWebAutoConfiguration { @Bean public SpanExtractor httpServletRequestSpanExtractor( - SkipPatternProvider skipPatternProvider, Random random) { - return new HttpServletRequestExtractor(skipPatternProvider.skipPattern(), random); + SkipPatternProvider skipPatternProvider) { + return new HttpServletRequestExtractor(skipPatternProvider.skipPattern()); } @Bean diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/SpanTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/SpanTests.java index 7a09ae389..c8eb2a071 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/SpanTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/SpanTests.java @@ -109,4 +109,9 @@ public class SpanTests { then(deserialized.tags()) .isEqualTo(span.tags()); } + + @Test(expected = IllegalArgumentException.class) + public void should_throw_exception_when_converting_invalid_hex_value() { + Span.hexToId("invalid"); + } } diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractorTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractorTests.java index 8da2be3f7..acd296951 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractorTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/messaging/MessagingSpanExtractorTests.java @@ -27,6 +27,7 @@ import org.springframework.messaging.MessageHeaders; import org.springframework.messaging.support.MessageBuilder; import org.springframework.util.StringUtils; +import static org.assertj.core.api.Assertions.fail; import static org.assertj.core.api.BDDAssertions.then; import static org.springframework.cloud.sleuth.assertions.SleuthAssertions.then; @@ -47,10 +48,12 @@ public class MessagingSpanExtractorTests { Message message = MessageBuilder.createMessage("", headers("invalid", randomId())); - Span span = this.extractor.joinTrace(message); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); + try { + this.extractor.joinTrace(message); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } @Test @@ -58,11 +61,12 @@ public class MessagingSpanExtractorTests { Message message = MessageBuilder.createMessage("", headers(randomId(), "invalid")); - Span span = this.extractor.joinTrace(message); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); - then(span.getSpanId()).isNotZero(); + try { + this.extractor.joinTrace(message); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } @Test @@ -70,12 +74,12 @@ public class MessagingSpanExtractorTests { Message message = MessageBuilder.createMessage("", headers(randomId(), randomId(), "invalid")); - Span span = this.extractor.joinTrace(message); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); - then(span.getSpanId()).isNotZero(); - then(span.getParents()).isEmpty(); + try { + this.extractor.joinTrace(message); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } private MessageHeaders headers() { 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 223e996b6..fc660f2ad 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 @@ -28,6 +28,7 @@ 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) @@ -35,7 +36,7 @@ public class HttpServletRequestExtractorTests { @Mock HttpServletRequest request; HttpServletRequestExtractor extractor = new HttpServletRequestExtractor( - Pattern.compile(""), new Random()); + Pattern.compile("")); @Before public void setup() { @@ -53,10 +54,12 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.TRACE_ID_NAME)) .willReturn("invalid"); - Span span = this.extractor.joinTrace(this.request); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); + try { + this.extractor.joinTrace(this.request); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } @Test @@ -66,11 +69,12 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.SPAN_ID_NAME)) .willReturn("invalid"); - Span span = this.extractor.joinTrace(this.request); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); - then(span.getSpanId()).isNotZero(); + try { + this.extractor.joinTrace(this.request); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } @Test @@ -82,11 +86,11 @@ public class HttpServletRequestExtractorTests { BDDMockito.given(this.request.getHeader(Span.PARENT_ID_NAME)) .willReturn("invalid"); - Span span = this.extractor.joinTrace(this.request); - - then(span).isNotNull(); - then(span.getTraceId()).isNotZero(); - then(span.getSpanId()).isNotZero(); - then(span.getParents()).isEmpty(); + try { + this.extractor.joinTrace(this.request); + fail("should throw an exception"); + } catch (IllegalArgumentException e) { + then(e).hasMessageContaining("Malformed id"); + } } } \ No newline at end of file diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterMockChainIntegrationTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterMockChainIntegrationTests.java index c62a7d1ca..f8dbb0134 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterMockChainIntegrationTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/instrument/web/TraceFilterMockChainIntegrationTests.java @@ -73,8 +73,7 @@ public class TraceFilterMockChainIntegrationTests { @Test public void startsNewTrace() throws Exception { TraceFilter filter = new TraceFilter(this.tracer, this.traceKeys, new NoOpSpanReporter(), - new HttpServletRequestExtractor(Pattern.compile(TraceFilter.DEFAULT_SKIP_PATTERN), - new Random()), + new HttpServletRequestExtractor(Pattern.compile(TraceFilter.DEFAULT_SKIP_PATTERN)), new HttpServletResponseInjector(), keysInjector); filter.doFilter(this.request, this.response, this.filterChain); assertNull(TestSpanContextHolder.getCurrentSpan()); @@ -86,8 +85,7 @@ public class TraceFilterMockChainIntegrationTests { this.request = builder().header(Span.SPAN_ID_NAME, generator.nextLong()) .header(Span.TRACE_ID_NAME, generator.nextLong()).buildRequest(new MockServletContext()); TraceFilter filter = new TraceFilter(this.tracer, this.traceKeys, new NoOpSpanReporter(), - new HttpServletRequestExtractor(Pattern.compile(TraceFilter.DEFAULT_SKIP_PATTERN), - new Random()), + new HttpServletRequestExtractor(Pattern.compile(TraceFilter.DEFAULT_SKIP_PATTERN)), new HttpServletResponseInjector(), keysInjector); filter.doFilter(this.request, this.response, this.filterChain); assertNull(TestSpanContextHolder.getCurrentSpan()); 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 9030e5d59..df9e1eac9 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 @@ -63,7 +63,7 @@ public class TraceFilterTests { @Mock SpanLogger spanLogger; ArrayListSpanAccumulator spanReporter = new ArrayListSpanAccumulator(); SpanExtractor spanExtractor = new HttpServletRequestExtractor(Pattern - .compile(TraceFilter.DEFAULT_SKIP_PATTERN), new Random()); + .compile(TraceFilter.DEFAULT_SKIP_PATTERN)); SpanInjector spanInjector = new HttpServletResponseInjector(); private Tracer tracer; @@ -317,6 +317,19 @@ public class TraceFilterTests { then(TestSpanContextHolder.getCurrentSpan()).isNull(); } + @Test + public void returns400IfSpanIsMalformed() 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, + this.spanExtractor, this.spanInjector, this.httpTraceKeysInjector); + + filter.doFilter(this.request, this.response, this.filterChain); + + then(TestSpanContextHolder.getCurrentSpan()).isNull(); + then(this.response.getStatus()).isEqualTo(HttpStatus.BAD_REQUEST.value()); + } + public void verifyParentSpanHttpTags() { verifyParentSpanHttpTags(HttpStatus.OK); }