From 7d7c68902608de81ac9c06e97bac0aaca0a5264f Mon Sep 17 00:00:00 2001 From: w1am <33353798+w1am@users.noreply.github.com> Date: Mon, 14 Sep 2026 21:09:39 +0400 Subject: [PATCH] fix: propagate W3C trace context for unsampled traces in event metadata --- .../io/kurrent/dbclient/ClientTelemetry.java | 93 +++++++++--- .../dbclient/ClientTelemetryConstants.java | 2 + .../java/io/kurrent/dbclient/MiscTests.java | 2 +- .../TracingContextPropagationTests.java | 139 ++++++++++++++++++ .../kurrent/dbclient/streams/AppendTests.java | 4 +- .../StreamsTracingInstrumentationTests.java | 2 + 6 files changed, 217 insertions(+), 25 deletions(-) create mode 100644 src/test/java/io/kurrent/dbclient/TracingContextPropagationTests.java diff --git a/src/main/java/io/kurrent/dbclient/ClientTelemetry.java b/src/main/java/io/kurrent/dbclient/ClientTelemetry.java index b4f1c4cc..1816c30b 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,36 @@ class ClientTelemetry { put(ClientTelemetryAttributes.Database.SYSTEM, ClientTelemetryConstants.INSTRUMENTATION_NAME); }}; + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + + private static final TextMapSetter METADATA_SETTER = + (userMetadata, key, value) -> userMetadata.put("$" + key, value); + + private static final TextMapGetter METADATA_GETTER = new TextMapGetter() { + @Override + public Iterable keys(ObjectNode userMetadata) { + return Arrays.asList("traceparent", "tracestate"); + } + + @Override + public String get(ObjectNode userMetadata, String key) { + return getTextField(userMetadata, "$" + key); + } + }; + + 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,49 +66,71 @@ private static List tryInjectTracingContext(Span span, List traceAppend( BiFunction, CompletableFuture> appendOperation, ManagedChannel channel, diff --git a/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java b/src/main/java/io/kurrent/dbclient/ClientTelemetryConstants.java index f4d77dc5..0ce2c1fd 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/MiscTests.java b/src/test/java/io/kurrent/dbclient/MiscTests.java index 658a2781..ab7d1264 100644 --- a/src/test/java/io/kurrent/dbclient/MiscTests.java +++ b/src/test/java/io/kurrent/dbclient/MiscTests.java @@ -6,5 +6,5 @@ @Suite @SelectPackages("io.kurrent.dbclient.misc") -@SelectClasses({SubscriptionStreamConsumerTests.class, LeaderRedirectUnitTest.class}) +@SelectClasses({SubscriptionStreamConsumerTests.class, LeaderRedirectUnitTest.class, TracingContextPropagationTests.class}) public class MiscTests {} 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..5380aa62 --- /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 class TracingContextPropagationTests { + private static final String TRACE_ID = "0af7651916cd43dd8448eb211c80319c"; + private static final String SPAN_ID = "b7ad6b7169203331"; + private static final String STALE_METADATA = "{" + + "\"$traceparent\":\"00-11111111111111111111111111111111-1111111111111111-01\"," + + "\"$tracestate\":\"dd=s:1\"," + + "\"$traceId\":\"11111111111111111111111111111111\"," + + "\"$spanId\":\"1111111111111111\"" + + "}"; + + private static Span spanWith(TraceFlags flags, TraceState traceState) { + return Span.wrap(SpanContext.create(TRACE_ID, SPAN_ID, flags, traceState)); + } + + private static ObjectNode parseMetadata(byte[] metadata) throws Exception { + return new ObjectMapper().readValue(metadata, ObjectNode.class); + } + + @Test + public 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 + public 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 + public 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 + public 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 + public 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 + public void testExtractionReturnsNullWhenNoTracingMetadataIsPresent() { + Assertions.assertNull(ClientTelemetry.tryExtractTracingContext(null)); + Assertions.assertNull(ClientTelemetry.tryExtractTracingContext( + "{\"foo\":\"bar\"}".getBytes(StandardCharsets.UTF_8))); + } + + @Test + public 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/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