Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
93 changes: 70 additions & 23 deletions src/main/java/io/kurrent/dbclient/ClientTelemetry.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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<ObjectNode> METADATA_SETTER =
(userMetadata, key, value) -> userMetadata.put("$" + key, value);

private static final TextMapGetter<ObjectNode> METADATA_GETTER = new TextMapGetter<ObjectNode>() {
@Override
public Iterable<String> 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<EventData> tryInjectTracingContext(Span span, List<EventData> events) {
if (!span.getSpanContext().isValid() || !span.getSpanContext().isSampled())
static List<EventData> tryInjectTracingContext(Span span, List<EventData> events) {
if (!span.getSpanContext().isValid())
return events;

List<EventData> injectedEvents = new ArrayList<>();
Expand All @@ -41,49 +66,71 @@ private static List<EventData> tryInjectTracingContext(Span span, List<EventData
return injectedEvents;
}

private static byte[] tryInjectTracingContext(Span span, byte[] userMetadataBytes) {
static byte[] tryInjectTracingContext(Span span, byte[] userMetadataBytes) {
if (!span.getSpanContext().isValid())
return userMetadataBytes;

try {
ObjectMapper objectMapper = new ObjectMapper();
ObjectNode userMetadata = userMetadataBytes != null
? objectMapper.readValue(userMetadataBytes, ObjectNode.class)
: objectMapper.createObjectNode();
? OBJECT_MAPPER.readValue(userMetadataBytes, ObjectNode.class)
: OBJECT_MAPPER.createObjectNode();

userMetadata.remove(ClientTelemetryConstants.Metadata.TRACE_STATE);

userMetadata.put(ClientTelemetryConstants.Metadata.TRACE_ID, span.getSpanContext().getTraceId());
userMetadata.put(ClientTelemetryConstants.Metadata.SPAN_ID, span.getSpanContext().getSpanId());
W3CTraceContextPropagator.getInstance()
.inject(Context.root().with(span), userMetadata, METADATA_SETTER);

return objectMapper.writeValueAsBytes(userMetadata);
// Legacy fields carry no sampling flag, so older readers would treat an
// unsampled trace as sampled. Only write them when sampled.
if (span.getSpanContext().isSampled()) {
userMetadata.put(ClientTelemetryConstants.Metadata.TRACE_ID, span.getSpanContext().getTraceId());
userMetadata.put(ClientTelemetryConstants.Metadata.SPAN_ID, span.getSpanContext().getSpanId());
} else {
userMetadata.remove(ClientTelemetryConstants.Metadata.TRACE_ID);
userMetadata.remove(ClientTelemetryConstants.Metadata.SPAN_ID);
}

return OBJECT_MAPPER.writeValueAsBytes(userMetadata);
} catch (Throwable t) {
// User metadata may not be a valid JSON object, or not JSON altogether.
return userMetadataBytes;
}
}

private static SpanContext tryExtractTracingContext(byte[] userMetadataBytes) {
static SpanContext tryExtractTracingContext(byte[] userMetadataBytes) {
if (userMetadataBytes == null)
return null;

try {
ObjectNode userMetadata = new ObjectMapper().readValue(userMetadataBytes, ObjectNode.class);

JsonNode traceIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.TRACE_ID);
JsonNode spanIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.SPAN_ID);

if (traceIdNode == null || spanIdNode == null)
return null;
ObjectNode userMetadata = OBJECT_MAPPER.readValue(userMetadataBytes, ObjectNode.class);

String traceId = traceIdNode.asText();
String spanId = spanIdNode.asText();
Context extracted = W3CTraceContextPropagator.getInstance()
.extract(Context.root(), userMetadata, METADATA_GETTER);

if (!TraceId.isValid(traceId) || !SpanId.isValid(spanId))
return null;
SpanContext traceParentContext = Span.fromContext(extracted).getSpanContext();
if (traceParentContext.isValid())
return traceParentContext;

return SpanContext.createFromRemoteParent(traceId, spanId, TraceFlags.getSampled(),
TraceState.getDefault());
return tryExtractLegacyTracingContext(userMetadata);
} catch (Throwable t) {
return null;
}
}

private static SpanContext tryExtractLegacyTracingContext(ObjectNode userMetadata) {
String traceId = getTextField(userMetadata, ClientTelemetryConstants.Metadata.TRACE_ID);
String spanId = getTextField(userMetadata, ClientTelemetryConstants.Metadata.SPAN_ID);

if (traceId == null || spanId == null)
return null;

if (!TraceId.isValid(traceId) || !SpanId.isValid(spanId))
return null;

return SpanContext.createFromRemoteParent(traceId, spanId, TraceFlags.getSampled(),
TraceState.getDefault());
}

static CompletableFuture<WriteResult> traceAppend(
BiFunction<ManagedChannel, List<EventData>, CompletableFuture<WriteResult>> appendOperation,
ManagedChannel channel,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down
2 changes: 1 addition & 1 deletion src/test/java/io/kurrent/dbclient/MiscTests.java
Original file line number Diff line number Diff line change
Expand Up @@ -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 {}
139 changes: 139 additions & 0 deletions src/test/java/io/kurrent/dbclient/TracingContextPropagationTests.java
Original file line number Diff line number Diff line change
@@ -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<EventData> 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"));
}
}
4 changes: 3 additions & 1 deletion src/test/java/io/kurrent/dbclient/streams/AppendTests.java
Original file line number Diff line number Diff line change
Expand Up @@ -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))
);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading