Throwing exception when tracing is malformed (#308)

with this change we will throw an IllegalArgumentException when the malformed tracing data are sent. In TraceFilter we're catching it and sending back 400 response. In case of messaging the exception gets propagated

fixed #306

* Changes following review

* Changes following review
This commit is contained in:
Marcin Grzejszczak
2016-06-22 14:09:12 +02:00
committed by GitHub
parent 6bea07af1e
commit 117a5e8347
10 changed files with 87 additions and 117 deletions

View File

@@ -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

View File

@@ -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<Message<?>> {
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<Message<?>> {
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<Message<?>> {
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));
}
}
}

View File

@@ -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<HttpServletRequest> {
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<HttpServletRequest> {
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<HttpServletRequest> {
+ "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<HttpServletRequest> {
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<HttpServletRequest> {
}
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);
}
}
}

View File

@@ -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);

View File

@@ -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<HttpServletRequest> httpServletRequestSpanExtractor(
SkipPatternProvider skipPatternProvider, Random random) {
return new HttpServletRequestExtractor(skipPatternProvider.skipPattern(), random);
SkipPatternProvider skipPatternProvider) {
return new HttpServletRequestExtractor(skipPatternProvider.skipPattern());
}
@Bean

View File

@@ -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");
}
}

View File

@@ -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() {

View File

@@ -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");
}
}
}

View File

@@ -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());

View File

@@ -63,7 +63,7 @@ public class TraceFilterTests {
@Mock SpanLogger spanLogger;
ArrayListSpanAccumulator spanReporter = new ArrayListSpanAccumulator();
SpanExtractor<HttpServletRequest> spanExtractor = new HttpServletRequestExtractor(Pattern
.compile(TraceFilter.DEFAULT_SKIP_PATTERN), new Random());
.compile(TraceFilter.DEFAULT_SKIP_PATTERN));
SpanInjector<HttpServletResponse> 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);
}