diff --git a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/DefaultTracerTests.java b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/DefaultTracerTests.java index 04fcccb56..fdafe1d97 100644 --- a/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/DefaultTracerTests.java +++ b/spring-cloud-sleuth-core/src/test/java/org/springframework/cloud/sleuth/DefaultTracerTests.java @@ -16,6 +16,7 @@ package org.springframework.cloud.sleuth; +import static org.hamcrest.CoreMatchers.equalTo; import static org.hamcrest.Matchers.is; import static org.junit.Assert.assertThat; import static org.mockito.Matchers.isA; @@ -123,6 +124,34 @@ public class DefaultTracerTests { assertThat(child.isExportable(), is(false)); } + @Test + public void parentNotRemovedIfActiveOnJoin() { + DefaultTracer tracer = new DefaultTracer(new AlwaysSampler(), new Random(), this.publisher); + Span parent = tracer.startTrace(CREATE_SIMPLE_TRACE); + Span span = tracer.joinTrace(IMPORTANT_WORK_1, parent); + tracer.close(span); + assertThat(tracer.getCurrentSpan(), is(equalTo(parent))); + } + + @Test + public void parentRemovedIfNotActiveOnJoin() { + DefaultTracer tracer = new DefaultTracer(new AlwaysSampler(), new Random(), this.publisher); + Span parent = Span.builder().name(CREATE_SIMPLE_TRACE).traceId(1L).spanId(1L).build(); + Span span = tracer.joinTrace(IMPORTANT_WORK_1, parent); + tracer.close(span); + assertThat(tracer.getCurrentSpan(), is(equalTo(null))); + } + + @Test + public void grandParentRestoredAfterAutoClose() { + DefaultTracer tracer = new DefaultTracer(new AlwaysSampler(), new Random(), this.publisher); + Span grandParent = tracer.startTrace(CREATE_SIMPLE_TRACE); + Span parent = Span.builder().name(IMPORTANT_WORK_1).traceId(1L).spanId(1L).build(); + Span span = tracer.joinTrace(IMPORTANT_WORK_2, parent); + tracer.close(span); + assertThat(tracer.getCurrentSpan(), is(equalTo(grandParent))); + } + private Span assertSpan(List spans, Long parentId, String name) { List found = findSpans(spans, parentId); assertThat("more than one span with parentId " + parentId, found.size(), is(1));