Skip to content
Draft
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
114 changes: 91 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,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<ObjectNode> 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<ObjectNode> METADATA_GETTER = new TextMapGetter<ObjectNode>() {
@Override
public Iterable<String> 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<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,47 +84,72 @@ 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);
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);
ObjectNode userMetadata = OBJECT_MAPPER.readValue(userMetadataBytes, ObjectNode.class);

JsonNode traceIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.TRACE_ID);
JsonNode spanIdNode = userMetadata.get(ClientTelemetryConstants.Metadata.SPAN_ID);
SpanContext traceParentContext = tryExtractTraceParentContext(userMetadata);
if (traceParentContext != null)
return traceParentContext;

if (traceIdNode == null || spanIdNode == null)
return null;
return tryExtractLegacyTracingContext(userMetadata);
} catch (Throwable t) {
return null;
}
}

String traceId = traceIdNode.asText();
String spanId = spanIdNode.asText();
private static SpanContext tryExtractTraceParentContext(ObjectNode userMetadata) {
Context extractedContext = W3CTraceContextPropagator.getInstance()
.extract(Context.root(), userMetadata, METADATA_GETTER);

if (!TraceId.isValid(traceId) || !SpanId.isValid(spanId))
return null;
SpanContext spanContext = Span.fromContext(extractedContext).getSpanContext();
return spanContext.isValid() ? spanContext : null;
}

return SpanContext.createFromRemoteParent(traceId, spanId, TraceFlags.getSampled(),
TraceState.getDefault());
} catch (Throwable t) {
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(
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/TelemetryTests.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down
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 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<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
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"));
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
package io.kurrent.dbclient;

public class TracingContextPropagationUnitTests implements TracingContextPropagationTests {
}
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