diff --git a/dd-java-agent/agent-llmobs/src/main/java/datadog/trace/llmobs/domain/DDLLMObsSpan.java b/dd-java-agent/agent-llmobs/src/main/java/datadog/trace/llmobs/domain/DDLLMObsSpan.java index 19a1c73346d..86913f110c7 100644 --- a/dd-java-agent/agent-llmobs/src/main/java/datadog/trace/llmobs/domain/DDLLMObsSpan.java +++ b/dd-java-agent/agent-llmobs/src/main/java/datadog/trace/llmobs/domain/DDLLMObsSpan.java @@ -156,35 +156,81 @@ public void annotateIO(List inputData, List inputDocuments, String outputData) { + if (finished) { + return; + } + if (inputDocuments != null && !inputDocuments.isEmpty()) { + span.setTag(INPUT, inputDocuments); + } + if (outputData != null && !outputData.isEmpty()) { + span.setTag(OUTPUT, outputData); + } + } + + @Override + public void annotateRetrievalIO(String inputData, List outputDocuments) { if (finished) { return; } - boolean wrongSpanKind = false; if (inputData != null && !inputData.isEmpty()) { - if (Tags.LLMOBS_LLM_SPAN_KIND.equals(spanKind)) { - wrongSpanKind = true; - annotateIO( - Collections.singletonList(LLMObs.LLMMessage.from(LLM_MESSAGE_UNKNOWN_ROLE, inputData)), - null); - } else { - span.setTag(INPUT, inputData); + span.setTag(INPUT, inputData); + } + if (outputDocuments != null && !outputDocuments.isEmpty()) { + span.setTag(OUTPUT, outputDocuments); + } + } + + @Override + public void annotateIO(String inputData, String outputData) { + if (finished) { + return; + } + boolean hasInput = inputData != null && !inputData.isEmpty(); + boolean hasOutput = outputData != null && !outputData.isEmpty(); + if (Tags.LLMOBS_LLM_SPAN_KIND.equals(spanKind)) { + List inputMessages = + hasInput + ? Collections.singletonList( + LLMObs.LLMMessage.from(LLM_MESSAGE_UNKNOWN_ROLE, inputData)) + : null; + List outputMessages = + hasOutput + ? Collections.singletonList( + LLMObs.LLMMessage.from(LLM_MESSAGE_UNKNOWN_ROLE, outputData)) + : null; + annotateIO(inputMessages, outputMessages); + if (hasInput || hasOutput) { + LOGGER.warn( + "the span being annotated is an LLM span, it is recommended to use the overload with List as arguments"); } + return; } - if (outputData != null && !outputData.isEmpty()) { - if (Tags.LLMOBS_LLM_SPAN_KIND.equals(spanKind)) { - wrongSpanKind = true; - annotateIO( - null, - Collections.singletonList( - LLMObs.LLMMessage.from(LLM_MESSAGE_UNKNOWN_ROLE, outputData))); - } else { - span.setTag(OUTPUT, outputData); + if (Tags.LLMOBS_EMBEDDING_SPAN_KIND.equals(spanKind)) { + List inputDocuments = + hasInput ? Collections.singletonList(LLMObs.Document.from(inputData)) : null; + annotateEmbeddingIO(inputDocuments, outputData); + if (hasInput) { + LOGGER.warn( + "the span being annotated is an embedding span, it is recommended to use annotateEmbeddingIO"); + } + return; + } + if (Tags.LLMOBS_RETRIEVAL_SPAN_KIND.equals(spanKind)) { + List outputDocuments = + hasOutput ? Collections.singletonList(LLMObs.Document.from(outputData)) : null; + annotateRetrievalIO(inputData, outputDocuments); + if (hasOutput) { + LOGGER.warn( + "the span being annotated is a retrieval span, it is recommended to use annotateRetrievalIO"); } + return; + } + if (hasInput) { + span.setTag(INPUT, inputData); } - if (wrongSpanKind) { - LOGGER.warn( - "the span being annotated is an LLM span, it is recommended to use the overload with List as arguments"); + if (hasOutput) { + span.setTag(OUTPUT, outputData); } } diff --git a/dd-java-agent/agent-llmobs/src/test/java/datadog/trace/llmobs/domain/DDLLMObsSpanDocumentIOTest.java b/dd-java-agent/agent-llmobs/src/test/java/datadog/trace/llmobs/domain/DDLLMObsSpanDocumentIOTest.java new file mode 100644 index 00000000000..10e8a6af54a --- /dev/null +++ b/dd-java-agent/agent-llmobs/src/test/java/datadog/trace/llmobs/domain/DDLLMObsSpanDocumentIOTest.java @@ -0,0 +1,141 @@ +package datadog.trace.llmobs.domain; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNull; + +import datadog.trace.agent.tooling.TracerInstaller; +import datadog.trace.api.WellKnownTags; +import datadog.trace.api.llmobs.LLMObs; +import datadog.trace.bootstrap.instrumentation.api.AgentSpan; +import datadog.trace.bootstrap.instrumentation.api.Tags; +import datadog.trace.core.CoreTracer; +import java.lang.reflect.Field; +import java.util.Arrays; +import java.util.List; +import org.junit.jupiter.api.AfterAll; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; + +class DDLLMObsSpanDocumentIOTest { + private static final String INPUT_TAG = "_ml_obs_tag.input"; + private static final String OUTPUT_TAG = "_ml_obs_tag.output"; + private static final Field SPAN_FIELD; + + private static CoreTracer tracer; + + static { + try { + SPAN_FIELD = DDLLMObsSpan.class.getDeclaredField("span"); + SPAN_FIELD.setAccessible(true); + } catch (ReflectiveOperationException error) { + throw new ExceptionInInitializerError(error); + } + } + + @BeforeAll + static void installTracer() { + tracer = CoreTracer.builder().build(); + TracerInstaller.forceInstallGlobalTracer(tracer); + } + + @AfterAll + static void closeTracer() { + TracerInstaller.forceInstallGlobalTracer(null); + tracer.close(); + } + + @Test + void formatsEmbeddingStringInputAsDocument() throws IllegalAccessException { + DDLLMObsSpan llmObsSpan = newSpan(Tags.LLMOBS_EMBEDDING_SPAN_KIND); + try { + llmObsSpan.annotateIO("embedding input", "embedding output"); + + AgentSpan span = (AgentSpan) SPAN_FIELD.get(llmObsSpan); + assertDocument(span.getTag(INPUT_TAG), "embedding input"); + assertEquals("embedding output", span.getTag(OUTPUT_TAG)); + } finally { + llmObsSpan.finish(); + } + } + + @Test + void formatsRetrievalStringOutputAsDocument() throws IllegalAccessException { + DDLLMObsSpan llmObsSpan = newSpan(Tags.LLMOBS_RETRIEVAL_SPAN_KIND); + try { + llmObsSpan.annotateIO("retrieval input", "retrieval output"); + + AgentSpan span = (AgentSpan) SPAN_FIELD.get(llmObsSpan); + assertEquals("retrieval input", span.getTag(INPUT_TAG)); + assertDocument(span.getTag(OUTPUT_TAG), "retrieval output"); + } finally { + llmObsSpan.finish(); + } + } + + @Test + void acceptsEmbeddingDocumentInputs() throws IllegalAccessException { + DDLLMObsSpan llmObsSpan = newSpan(Tags.LLMOBS_EMBEDDING_SPAN_KIND); + List documents = + Arrays.asList( + LLMObs.Document.from("first input", "first.txt", "input-1", 0.5), + LLMObs.Document.from("second input")); + try { + llmObsSpan.annotateEmbeddingIO(documents, "embedding output"); + + AgentSpan span = (AgentSpan) SPAN_FIELD.get(llmObsSpan); + assertEquals(documents, span.getTag(INPUT_TAG)); + assertDocument(documents.get(0), "first input", "first.txt", "input-1", 0.5); + assertEquals("embedding output", span.getTag(OUTPUT_TAG)); + } finally { + llmObsSpan.finish(); + } + } + + @Test + void acceptsRetrievalDocumentOutputs() throws IllegalAccessException { + DDLLMObsSpan llmObsSpan = newSpan(Tags.LLMOBS_RETRIEVAL_SPAN_KIND); + List documents = + Arrays.asList( + LLMObs.Document.from("first output", "result.txt", "output-1", 0.95), + LLMObs.Document.from("second output")); + try { + llmObsSpan.annotateRetrievalIO("retrieval input", documents); + + AgentSpan span = (AgentSpan) SPAN_FIELD.get(llmObsSpan); + assertEquals("retrieval input", span.getTag(INPUT_TAG)); + assertEquals(documents, span.getTag(OUTPUT_TAG)); + assertDocument(documents.get(0), "first output", "result.txt", "output-1", 0.95); + } finally { + llmObsSpan.finish(); + } + } + + private static DDLLMObsSpan newSpan(String kind) { + WellKnownTags tags = + new WellKnownTags("runtime-id", "hostname", "test", "service", "version", "java"); + return new DDLLMObsSpan(kind, "span", "ml-app", null, "service", tags); + } + + private static void assertDocument(Object value, String expectedText) { + List documents = assertInstanceOf(List.class, value); + assertEquals(1, documents.size()); + LLMObs.Document document = assertInstanceOf(LLMObs.Document.class, documents.get(0)); + assertEquals(expectedText, document.getText()); + assertNull(document.getName()); + assertNull(document.getId()); + assertNull(document.getScore()); + } + + private static void assertDocument( + LLMObs.Document document, + String expectedText, + String expectedName, + String expectedId, + double expectedScore) { + assertEquals(expectedText, document.getText()); + assertEquals(expectedName, document.getName()); + assertEquals(expectedId, document.getId()); + assertEquals(expectedScore, document.getScore()); + } +} diff --git a/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObs.java b/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObs.java index 25f3ff0a8ac..522e871e98c 100644 --- a/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObs.java +++ b/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObs.java @@ -305,17 +305,40 @@ public List getToolResults() { public static class Document { private String text; + private String name; + private String id; + private Double score; public static Document from(String text) { - return new Document(text); + return new Document(text, null, null, null); + } + + public static Document from( + String text, @Nullable String name, @Nullable String id, @Nullable Double score) { + return new Document(text, name, id, score); } - private Document(String text) { + private Document(String text, String name, String id, Double score) { this.text = text; + this.name = name; + this.id = id; + this.score = score; } public String getText() { return text; } + + public String getName() { + return name; + } + + public String getId() { + return id; + } + + public Double getScore() { + return score; + } } } diff --git a/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObsSpan.java b/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObsSpan.java index 9d778da56cf..817cfc96e85 100644 --- a/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObsSpan.java +++ b/dd-trace-api/src/main/java/datadog/trace/api/llmobs/LLMObsSpan.java @@ -15,6 +15,26 @@ public interface LLMObsSpan { */ void annotateIO(List inputMessages, List outputMessages); + /** + * Annotate an embedding span with document inputs and a string output. + * + * @param inputDocuments The input documents of the span + * @param outputData The output data of the span in the form of a string + */ + default void annotateEmbeddingIO(List inputDocuments, String outputData) { + annotateIO((String) null, outputData); + } + + /** + * Annotate a retrieval span with a string input and document outputs. + * + * @param inputData The input data of the span in the form of a string + * @param outputDocuments The output documents of the span + */ + default void annotateRetrievalIO(String inputData, List outputDocuments) { + annotateIO(inputData, (String) null); + } + /** * Annotate the span with inputs and outputs * diff --git a/dd-trace-core/src/main/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapper.java b/dd-trace-core/src/main/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapper.java index 155f4a59406..d55935e9835 100644 --- a/dd-trace-core/src/main/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapper.java +++ b/dd-trace-core/src/main/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapper.java @@ -357,6 +357,9 @@ public void accept(Metadata metadata) { String key = tag.getKey().substring(LLMOBS_TAG_PREFIX.length()); Object val = tag.getValue(); if (key.equals(INPUT) || key.equals(OUTPUT)) { + boolean isDocumentIO = + (spanKind.equals(Tags.LLMOBS_EMBEDDING_SPAN_KIND) && key.equals(INPUT)) + || (spanKind.equals(Tags.LLMOBS_RETRIEVAL_SPAN_KIND) && key.equals(OUTPUT)); if (spanKind.equals(Tags.LLMOBS_LLM_SPAN_KIND)) { writable.writeString(key, null); if (val instanceof List) { @@ -371,24 +374,40 @@ public void accept(Metadata metadata) { val.getClass().getName()); continue; } - } else if (spanKind.equals(Tags.LLMOBS_EMBEDDING_SPAN_KIND) && key.equals(INPUT)) { - if (!(val instanceof List)) { - LOGGER.warn( - "unexpectedly found incorrect type for embedding span input {}, expecting list", - val.getClass().getName()); - continue; - } + } else if (isDocumentIO && isDocumentList(val)) { writable.writeString(key, null); writable.startMap(1); List documents = (List) val; writable.writeString("documents", null); writable.startArray(documents.size()); for (LLMObs.Document document : documents) { - writable.startMap(1); + int documentSize = 1; + if (document.getName() != null) documentSize++; + if (document.getId() != null) documentSize++; + if (document.getScore() != null) documentSize++; + writable.startMap(documentSize); writable.writeString("text", null); writable.writeString(document.getText(), null); + if (document.getName() != null) { + writable.writeString("name", null); + writable.writeString(document.getName(), null); + } + if (document.getId() != null) { + writable.writeString("id", null); + writable.writeString(document.getId(), null); + } + if (document.getScore() != null) { + writable.writeString("score", null); + writable.writeObject(document.getScore(), null); + } } } else { + if (isDocumentIO) { + LOGGER.warn( + "unexpectedly found invalid document data for {} span {}, serializing as value", + spanKind, + key); + } writable.writeString(key, null); writable.startMap(1); writable.writeString("value", null); @@ -444,6 +463,18 @@ private void writeToolDefinitions(List toolDefinitions) { } } + private static boolean isDocumentList(Object value) { + if (!(value instanceof List)) { + return false; + } + for (Object item : (List) value) { + if (!(item instanceof LLMObs.Document)) { + return false; + } + } + return true; + } + private void writeLlmInputMap(Map inputMap) { writable.startMap(inputMap.size()); for (Map.Entry entry : inputMap.entrySet()) { diff --git a/dd-trace-core/src/test/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapperTest.java b/dd-trace-core/src/test/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapperTest.java index 1da6bf4018f..a07ca4b6210 100644 --- a/dd-trace-core/src/test/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapperTest.java +++ b/dd-trace-core/src/test/java/datadog/trace/llmobs/writer/ddintake/LLMObsSpanMapperTest.java @@ -384,6 +384,109 @@ void testLLMObsSpanMapperOmitsTopLevelSessionIdWhenNotSet() throws Exception { tracer.close(); } + @Test + void testLLMObsSpanMapperSerializesDocumentIO() throws Exception { + LLMObsSpanMapper mapper = new LLMObsSpanMapper(); + CoreTracer tracer = tracerBuilder().writer(new ListWriter()).build(); + + AgentSpan embeddingSpan = + tracer + .buildSpan("datadog", "embedding") + .withTag("_ml_obs_tag.span.kind", Tags.LLMOBS_EMBEDDING_SPAN_KIND) + .withTag( + "_ml_obs_tag.input", + Arrays.asList( + LLMObs.Document.from("embedding input", "embedding.txt", "embedding-1", 0.75), + LLMObs.Document.from("embedding text only"))) + .start(); + embeddingSpan.setSpanType(InternalSpanTypes.LLMOBS); + embeddingSpan.finish(); + + AgentSpan retrievalSpan = + tracer + .buildSpan("datadog", "retrieval") + .withTag("_ml_obs_tag.span.kind", Tags.LLMOBS_RETRIEVAL_SPAN_KIND) + .withTag( + "_ml_obs_tag.output", + Collections.singletonList( + LLMObs.Document.from("retrieval output", "retrieval.txt", "retrieval-1", 0.9))) + .start(); + retrievalSpan.setSpanType(InternalSpanTypes.LLMOBS); + retrievalSpan.finish(); + + List nonDocumentInput = Arrays.asList("first raw input", "second raw input"); + AgentSpan embeddingWithNonDocumentInput = + tracer + .buildSpan("datadog", "embedding") + .withTag("_ml_obs_tag.span.kind", Tags.LLMOBS_EMBEDDING_SPAN_KIND) + .withTag("_ml_obs_tag.input", nonDocumentInput) + .start(); + embeddingWithNonDocumentInput.setSpanType(InternalSpanTypes.LLMOBS); + embeddingWithNonDocumentInput.finish(); + + List nonDocumentOutput = Arrays.asList("first raw result", "second raw result"); + AgentSpan retrievalWithNonDocumentOutput = + tracer + .buildSpan("datadog", "retrieval") + .withTag("_ml_obs_tag.span.kind", Tags.LLMOBS_RETRIEVAL_SPAN_KIND) + .withTag("_ml_obs_tag.output", nonDocumentOutput) + .start(); + retrievalWithNonDocumentOutput.setSpanType(InternalSpanTypes.LLMOBS); + retrievalWithNonDocumentOutput.finish(); + + List trace = + Arrays.asList( + (DDSpan) embeddingSpan, + (DDSpan) retrievalSpan, + (DDSpan) embeddingWithNonDocumentInput, + (DDSpan) retrievalWithNonDocumentOutput); + CapturingByteBufferConsumer sink = new CapturingByteBufferConsumer(); + MsgPackWriter packer = new MsgPackWriter(new FlushingBuffer(16 * 1024, sink)); + packer.format(trace, mapper); + packer.flush(); + + datadog.trace.common.writer.Payload payload = mapper.newPayload(); + payload.withBody(trace.size(), sink.captured); + Map result = objectMapper.readValue(writeTo(payload), Map.class); + List> spans = (List>) result.get("spans"); + + Map embeddingMeta = (Map) spans.get(0).get("meta"); + Map embeddingInput = (Map) embeddingMeta.get("input"); + List> embeddingDocuments = + (List>) embeddingInput.get("documents"); + assertEquals("embedding input", embeddingDocuments.get(0).get("text")); + assertEquals("embedding.txt", embeddingDocuments.get(0).get("name")); + assertEquals("embedding-1", embeddingDocuments.get(0).get("id")); + assertEquals(0.75, embeddingDocuments.get(0).get("score")); + assertEquals("embedding text only", embeddingDocuments.get(1).get("text")); + assertFalse(embeddingDocuments.get(1).containsKey("name")); + assertFalse(embeddingDocuments.get(1).containsKey("id")); + assertFalse(embeddingDocuments.get(1).containsKey("score")); + assertFalse(embeddingInput.containsKey("value")); + + Map retrievalMeta = (Map) spans.get(1).get("meta"); + Map retrievalOutput = (Map) retrievalMeta.get("output"); + List> retrievalDocuments = + (List>) retrievalOutput.get("documents"); + assertEquals("retrieval output", retrievalDocuments.get(0).get("text")); + assertEquals("retrieval.txt", retrievalDocuments.get(0).get("name")); + assertEquals("retrieval-1", retrievalDocuments.get(0).get("id")); + assertEquals(0.9, retrievalDocuments.get(0).get("score")); + assertFalse(retrievalOutput.containsKey("value")); + + Map fallbackEmbeddingMeta = (Map) spans.get(2).get("meta"); + Map fallbackInput = (Map) fallbackEmbeddingMeta.get("input"); + assertEquals(nonDocumentInput, fallbackInput.get("value")); + assertFalse(fallbackInput.containsKey("documents")); + + Map fallbackRetrievalMeta = (Map) spans.get(3).get("meta"); + Map fallbackOutput = (Map) fallbackRetrievalMeta.get("output"); + assertEquals(nonDocumentOutput, fallbackOutput.get("value")); + assertFalse(fallbackOutput.containsKey("documents")); + + tracer.close(); + } + private static byte[] writeTo(datadog.trace.common.writer.Payload payload) throws IOException { ByteArrayOutputStream channel = new ByteArrayOutputStream(); payload.writeTo(