diff --git a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValve.java b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValve.java index 69387a122..c1cd6dad2 100644 --- a/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValve.java +++ b/spring-cloud-sleuth-instrumentation/src/main/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValve.java @@ -73,6 +73,19 @@ public class TraceValve extends ValveBase { @Override public void invoke(Request request, Response response) throws IOException, ServletException { + Object attribute = request.getAttribute(Span.class.getName()); + if (attribute != null) { + // this could happen for async dispatch + try (CurrentTraceContext.Scope ws = currentTraceContext().maybeScope(((Span) attribute).context())) { + Valve next = getNext(); + if (null == next) { + // no next valve + return; + } + next.invoke(request, response); + return; + } + } Exception ex = null; Span handleReceive = httpServerHandler().handleReceive(HttpServletRequestWrapper.create(request.getRequest())); if (log.isDebugEnabled()) { diff --git a/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValveTests.java b/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValveTests.java index 3d2e4c555..4366eb466 100644 --- a/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValveTests.java +++ b/spring-cloud-sleuth-instrumentation/src/test/java/org/springframework/cloud/sleuth/instrument/web/tomcat/TraceValveTests.java @@ -17,6 +17,7 @@ package org.springframework.cloud.sleuth.instrument.web.tomcat; import java.io.IOException; +import java.util.concurrent.atomic.AtomicInteger; import javax.servlet.ServletException; @@ -42,14 +43,21 @@ class TraceValveTests { SimpleSpan simpleSpan = new SimpleSpan(); + AtomicInteger startCounter = new AtomicInteger(); + + AtomicInteger endCounter = new AtomicInteger(); + HttpServerHandler httpServerHandler = new HttpServerHandler() { + @Override public SimpleSpan handleReceive(HttpServerRequest request) { + startCounter.incrementAndGet(); return simpleSpan.start(); } @Override public void handleSend(HttpServerResponse response, Span span) { + endCounter.incrementAndGet(); span.end(); } }; @@ -75,7 +83,9 @@ class TraceValveTests { private void thenSpanIsStartedAndStopped() { then(simpleSpan.started).isTrue(); + then(startCounter.get()).isEqualTo(1); then(simpleSpan.ended).isTrue(); + then(endCounter.get()).isEqualTo(1); } @Test @@ -94,6 +104,21 @@ class TraceValveTests { thenSpanIsStartedAndStopped(); } + @Test + void should_not_generate_a_new_span_when_one_already_present() throws ServletException, IOException { + Request request = request(); + + new TraceValve(this.httpServerHandler, new SimpleCurrentTraceContext()) { + @Override + public Valve getNext() { + return new TraceValve(httpServerHandler, new SimpleCurrentTraceContext()); + } + }.invoke(request, new Response()); + + then(request.getAttribute(TraceContext.class.getName())).isNotNull(); + thenSpanIsStartedAndStopped(); + } + private Request request() { Request request = new Request(new Connector()); request.setCoyoteRequest(new org.apache.coyote.Request()); diff --git a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java index 848a2416a..a79e151e0 100644 --- a/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java +++ b/tests/common/src/main/java/org/springframework/cloud/sleuth/instrument/messaging/TraceFunctionAroundWrapperTests.java @@ -47,7 +47,7 @@ public abstract class TraceFunctionAroundWrapperTests { public void test_tracing_with_supplier() { try (ConfigurableApplicationContext context = new SpringApplicationBuilder(configuration(), SampleConfiguration.class).run("--logging.level.org.springframework.cloud.function=DEBUG", - "--spring.main.lazy-initialization=true");) { + "--spring.main.lazy-initialization=true", "--server.port=0");) { TestSpanHandler spanHandler = context.getBean(TestSpanHandler.class); assertThat(spanHandler.reportedSpans()).isEmpty(); FunctionCatalog catalog = context.getBean(FunctionCatalog.class);