diff --git a/src/main/java/io/kurrent/dbclient/ClientTelemetry.java b/src/main/java/io/kurrent/dbclient/ClientTelemetry.java index b4f1c4cc..e00eec71 100644 --- a/src/main/java/io/kurrent/dbclient/ClientTelemetry.java +++ b/src/main/java/io/kurrent/dbclient/ClientTelemetry.java @@ -6,8 +6,11 @@ import io.grpc.ManagedChannel; import io.opentelemetry.api.GlobalOpenTelemetry; import io.opentelemetry.api.trace.*; +import io.opentelemetry.api.trace.propagation.W3CTraceContextPropagator; import io.opentelemetry.context.Context; import io.opentelemetry.context.Scope; +import io.opentelemetry.context.propagation.TextMapGetter; +import io.opentelemetry.context.propagation.TextMapSetter; import java.util.*; import java.util.concurrent.CompletableFuture; @@ -19,14 +22,54 @@ class ClientTelemetry { put(ClientTelemetryAttributes.Database.SYSTEM, ClientTelemetryConstants.INSTRUMENTATION_NAME); }}; + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + + private static final String W3C_TRACE_PARENT_KEY = "traceparent"; + private static final String W3C_TRACE_STATE_KEY = "tracestate"; + + private static final TextMapSetter METADATA_SETTER = (userMetadata, key, value) -> { + if (userMetadata == null) + return; + + if (W3C_TRACE_PARENT_KEY.equals(key)) + userMetadata.put(ClientTelemetryConstants.Metadata.TRACE_PARENT, value); + else if (W3C_TRACE_STATE_KEY.equals(key)) + userMetadata.put(ClientTelemetryConstants.Metadata.TRACE_STATE, value); + }; + + private static final TextMapGetter METADATA_GETTER = new TextMapGetter() { + @Override + public Iterable keys(ObjectNode userMetadata) { + return Arrays.asList(W3C_TRACE_PARENT_KEY, W3C_TRACE_STATE_KEY); + } + + @Override + public String get(ObjectNode userMetadata, String key) { + if (userMetadata == null) + return null; + + if (W3C_TRACE_PARENT_KEY.equals(key)) + return getTextField(userMetadata, ClientTelemetryConstants.Metadata.TRACE_PARENT); + if (W3C_TRACE_STATE_KEY.equals(key)) + return getTextField(userMetadata, ClientTelemetryConstants.Metadata.TRACE_STATE); + + return null; + } + }; + + private static String getTextField(ObjectNode userMetadata, String fieldName) { + JsonNode field = userMetadata.get(fieldName); + return field != null && field.isTextual() ? field.asText() : null; + } + private static Tracer getTracer() { return GlobalOpenTelemetry.getTracer( ClientTelemetry.class.getPackage().getName(), ClientTelemetry.class.getPackage().getImplementationVersion()); } - private static List tryInjectTracingContext(Span span, List events) { - if (!span.getSpanContext().isValid() || !span.getSpanContext().isSampled()) + static List tryInjectTracingContext(Span span, List events) { + if (!span.getSpanContext().isValid()) return events; List injectedEvents = new ArrayList<>(); @@ -41,47 +84,72 @@ private static List tryInjectTracingContext(Span span, List traceAppend( diff --git a/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java b/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java index f4d77dc5..71167f7a 100644 --- a/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java +++ b/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java @@ -6,6 +6,8 @@ public class ClientTelemetryConstants { public static class Metadata { public static final String TRACE_ID = "$traceId"; public static final String SPAN_ID = "$spanId"; + public static final String TRACE_PARENT = "$traceParent"; + public static final String TRACE_STATE = "$traceState"; } public static class Operations { diff --git a/src/test/java/io/kurrent/dbclient/TelemetryTests.java b/src/test/java/io/kurrent/dbclient/TelemetryTests.java index 5bfc0151..a385da14 100644 --- a/src/test/java/io/kurrent/dbclient/TelemetryTests.java +++ b/src/test/java/io/kurrent/dbclient/TelemetryTests.java @@ -26,7 +26,7 @@ import static io.opentelemetry.semconv.ServiceAttributes.SERVICE_NAME; -public class TelemetryTests implements StreamsTracingInstrumentationTests, PersistentSubscriptionsTracingInstrumentationTests, TracingContextInjectionTests { +public class TelemetryTests implements StreamsTracingInstrumentationTests, PersistentSubscriptionsTracingInstrumentationTests, TracingContextInjectionTests, TracingContextPropagationTests { static private Database database; static private Logger logger; diff --git a/src/test/java/io/kurrent/dbclient/TracingContextPropagationTests.java b/src/test/java/io/kurrent/dbclient/TracingContextPropagationTests.java new file mode 100644 index 00000000..724f96b4 --- /dev/null +++ b/src/test/java/io/kurrent/dbclient/TracingContextPropagationTests.java @@ -0,0 +1,139 @@ +package io.kurrent.dbclient; + +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.ObjectNode; +import io.opentelemetry.api.trace.Span; +import io.opentelemetry.api.trace.SpanContext; +import io.opentelemetry.api.trace.TraceFlags; +import io.opentelemetry.api.trace.TraceState; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +import java.nio.charset.StandardCharsets; +import java.util.Collections; +import java.util.List; + +public interface TracingContextPropagationTests { + String TRACE_ID = "0af7651916cd43dd8448eb211c80319c"; + String SPAN_ID = "b7ad6b7169203331"; + String STALE_METADATA = "{" + + "\"$traceParent\":\"00-11111111111111111111111111111111-1111111111111111-01\"," + + "\"$traceState\":\"dd=s:1\"," + + "\"$traceId\":\"11111111111111111111111111111111\"," + + "\"$spanId\":\"1111111111111111\"" + + "}"; + + default Span spanWith(TraceFlags flags, TraceState traceState) { + return Span.wrap(SpanContext.create(TRACE_ID, SPAN_ID, flags, traceState)); + } + + default ObjectNode parseMetadata(byte[] metadata) throws Exception { + return new ObjectMapper().readValue(metadata, ObjectNode.class); + } + + @Test + default void testInjectsSampledTraceContextAlongsideLegacyFields() throws Exception { + TraceState traceState = TraceState.builder().put("dd", "s:1").build(); + Span span = spanWith(TraceFlags.getSampled(), traceState); + byte[] userMetadata = "{\"foo\":\"bar\"}".getBytes(StandardCharsets.UTF_8); + + ObjectNode metadata = parseMetadata(ClientTelemetry.tryInjectTracingContext(span, userMetadata)); + + Assertions.assertEquals( + "00-" + TRACE_ID + "-" + SPAN_ID + "-01", + metadata.get(ClientTelemetryConstants.Metadata.TRACE_PARENT).asText()); + Assertions.assertEquals("dd=s:1", metadata.get(ClientTelemetryConstants.Metadata.TRACE_STATE).asText()); + Assertions.assertEquals(TRACE_ID, metadata.get(ClientTelemetryConstants.Metadata.TRACE_ID).asText()); + Assertions.assertEquals(SPAN_ID, metadata.get(ClientTelemetryConstants.Metadata.SPAN_ID).asText()); + Assertions.assertEquals("bar", metadata.get("foo").asText()); + } + + @Test + default void testInjectsUnsampledTraceContextAndStripsStaleTracingFields() throws Exception { + Span span = spanWith(TraceFlags.getDefault(), TraceState.getDefault()); + + ObjectNode metadata = parseMetadata(ClientTelemetry.tryInjectTracingContext( + span, STALE_METADATA.getBytes(StandardCharsets.UTF_8))); + + Assertions.assertEquals( + "00-" + TRACE_ID + "-" + SPAN_ID + "-00", + metadata.get(ClientTelemetryConstants.Metadata.TRACE_PARENT).asText()); + Assertions.assertNull(metadata.get(ClientTelemetryConstants.Metadata.TRACE_ID)); + Assertions.assertNull(metadata.get(ClientTelemetryConstants.Metadata.SPAN_ID)); + Assertions.assertNull(metadata.get(ClientTelemetryConstants.Metadata.TRACE_STATE)); + } + + @Test + default void testSkipsInjectionForInvalidSpanOrNonJsonObjectMetadata() { + List events = Collections.singletonList( + EventData.builderAsJson("TestEvent", "{}".getBytes(StandardCharsets.UTF_8)).build()); + byte[] jsonMetadata = "{\"foo\":\"bar\"}".getBytes(StandardCharsets.UTF_8); + byte[] nonJsonMetadata = "clearlynotvalidjson".getBytes(StandardCharsets.UTF_8); + Span validSpan = spanWith(TraceFlags.getSampled(), TraceState.getDefault()); + + Assertions.assertSame(events, ClientTelemetry.tryInjectTracingContext(Span.getInvalid(), events)); + Assertions.assertSame(jsonMetadata, ClientTelemetry.tryInjectTracingContext(Span.getInvalid(), jsonMetadata)); + Assertions.assertArrayEquals(nonJsonMetadata, ClientTelemetry.tryInjectTracingContext(validSpan, nonJsonMetadata)); + } + + @Test + default void testExtractionPrefersTraceParentAndPreservesFlagsAndTraceState() { + String metadata = "{" + + "\"$traceParent\":\"00-" + TRACE_ID + "-" + SPAN_ID + "-00\"," + + "\"$traceState\":\"dd=s:1\"," + + "\"$traceId\":\"11111111111111111111111111111111\"," + + "\"$spanId\":\"1111111111111111\"" + + "}"; + + SpanContext extracted = ClientTelemetry.tryExtractTracingContext(metadata.getBytes(StandardCharsets.UTF_8)); + + Assertions.assertNotNull(extracted); + Assertions.assertEquals(TRACE_ID, extracted.getTraceId()); + Assertions.assertEquals(SPAN_ID, extracted.getSpanId()); + Assertions.assertFalse(extracted.isSampled()); + Assertions.assertTrue(extracted.isRemote()); + Assertions.assertEquals("s:1", extracted.getTraceState().get("dd")); + } + + @Test + default void testExtractionFallsBackToLegacyFieldsAsSampled() { + String legacyOnly = "{\"$traceId\":\"" + TRACE_ID + "\",\"$spanId\":\"" + SPAN_ID + "\"}"; + String malformedTraceParent = "{" + + "\"$traceParent\":\"not-a-traceparent\"," + + "\"$traceId\":\"" + TRACE_ID + "\"," + + "\"$spanId\":\"" + SPAN_ID + "\"" + + "}"; + + for (String metadata : new String[]{legacyOnly, malformedTraceParent}) { + SpanContext extracted = ClientTelemetry.tryExtractTracingContext(metadata.getBytes(StandardCharsets.UTF_8)); + + Assertions.assertNotNull(extracted); + Assertions.assertEquals(TRACE_ID, extracted.getTraceId()); + Assertions.assertEquals(SPAN_ID, extracted.getSpanId()); + Assertions.assertTrue(extracted.isSampled()); + Assertions.assertTrue(extracted.isRemote()); + } + } + + @Test + default void testExtractionReturnsNullWhenNoTracingMetadataIsPresent() { + Assertions.assertNull(ClientTelemetry.tryExtractTracingContext(null)); + Assertions.assertNull(ClientTelemetry.tryExtractTracingContext( + "{\"foo\":\"bar\"}".getBytes(StandardCharsets.UTF_8))); + } + + @Test + default void testRoundTripPreservesSamplingDecisionAndTraceState() { + TraceState traceState = TraceState.builder().put("dd", "s:0").build(); + Span span = spanWith(TraceFlags.getDefault(), traceState); + + byte[] metadata = ClientTelemetry.tryInjectTracingContext(span, (byte[]) null); + SpanContext extracted = ClientTelemetry.tryExtractTracingContext(metadata); + + Assertions.assertNotNull(extracted); + Assertions.assertEquals(TRACE_ID, extracted.getTraceId()); + Assertions.assertEquals(SPAN_ID, extracted.getSpanId()); + Assertions.assertFalse(extracted.isSampled()); + Assertions.assertEquals("s:0", extracted.getTraceState().get("dd")); + } +} diff --git a/src/test/java/io/kurrent/dbclient/TracingContextPropagationUnitTests.java b/src/test/java/io/kurrent/dbclient/TracingContextPropagationUnitTests.java new file mode 100644 index 00000000..1966c047 --- /dev/null +++ b/src/test/java/io/kurrent/dbclient/TracingContextPropagationUnitTests.java @@ -0,0 +1,4 @@ +package io.kurrent.dbclient; + +public class TracingContextPropagationUnitTests implements TracingContextPropagationTests { +} diff --git a/src/test/java/io/kurrent/dbclient/streams/AppendTests.java b/src/test/java/io/kurrent/dbclient/streams/AppendTests.java index b2bc0b80..d66e2bfe 100644 --- a/src/test/java/io/kurrent/dbclient/streams/AppendTests.java +++ b/src/test/java/io/kurrent/dbclient/streams/AppendTests.java @@ -37,7 +37,9 @@ default void testAppendSingleEventNoStream() throws Throwable { () -> Assertions.assertEquals(foo, mapper.readValue(first.getEventData(), Foo.class)), () -> Assertions.assertEquals(foo, mapper.readValue(first.getUserMetadata(), Foo.class)), () -> Assertions.assertFalse(userMetadata.has(ClientTelemetryConstants.Metadata.TRACE_ID)), - () -> Assertions.assertFalse(userMetadata.has(ClientTelemetryConstants.Metadata.SPAN_ID)) + () -> Assertions.assertFalse(userMetadata.has(ClientTelemetryConstants.Metadata.SPAN_ID)), + () -> Assertions.assertFalse(userMetadata.has(ClientTelemetryConstants.Metadata.TRACE_PARENT)), + () -> Assertions.assertFalse(userMetadata.has(ClientTelemetryConstants.Metadata.TRACE_STATE)) ); } diff --git a/src/test/java/io/kurrent/dbclient/telemetry/StreamsTracingInstrumentationTests.java b/src/test/java/io/kurrent/dbclient/telemetry/StreamsTracingInstrumentationTests.java index 4e91ed0f..c7f54955 100644 --- a/src/test/java/io/kurrent/dbclient/telemetry/StreamsTracingInstrumentationTests.java +++ b/src/test/java/io/kurrent/dbclient/telemetry/StreamsTracingInstrumentationTests.java @@ -60,9 +60,11 @@ default void testTracingContextIsInjectedAsExpectedWhenUserMetadataIsJsonObject( JsonNode traceIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.TRACE_ID); JsonNode spanIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.SPAN_ID); + JsonNode traceParentNode = userMetadata.get(ClientTelemetryConstants.Metadata.TRACE_PARENT); Assertions.assertNotNull(traceIdNode); Assertions.assertNotNull(spanIdNode); + Assertions.assertNotNull(traceParentNode); } @Test