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
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,7 @@ private void tagSpan(
@Nullable String responseBody) {
try {
Map<String, Object> metadata = new java.util.HashMap<>();
metadata.put("provider", "gemini");
metadata.put("provider", "google");

// Parse request
if (requestBody != null) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -67,7 +67,7 @@ void testWrapGemini() {
span.getAttributes().get(AttributeKey.stringKey("braintrust.metadata"));
assertNotNull(metadataJson, "braintrust.metadata should be set");
var metadata = JSON_MAPPER.readTree(metadataJson);
assertEquals("gemini", metadata.get("provider").asText());
assertEquals("google", metadata.get("provider").asText());
assertEquals(MODEL_ID, metadata.get("model").asText());
assertEquals(0.0, metadata.get("temperature").asDouble());
assertEquals(50, metadata.get("maxOutputTokens").asInt());
Expand Down Expand Up @@ -145,7 +145,7 @@ void testWrapGeminiAsync() {
span.getAttributes().get(AttributeKey.stringKey("braintrust.metadata"));
assertNotNull(metadataJson, "braintrust.metadata should be set");
var metadata = JSON_MAPPER.readTree(metadataJson);
assertEquals("gemini", metadata.get("provider").asText());
assertEquals("google", metadata.get("provider").asText());
assertEquals(MODEL_ID, metadata.get("model").asText());
assertEquals(0.0, metadata.get("temperature").asDouble());
assertEquals(50, metadata.get("maxOutputTokens").asInt());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
import com.openai.models.responses.ResponseStreamEvent;
import dev.braintrust.bootstrap.BraintrustBridge;
import dev.braintrust.instrumentation.InstrumentationSemConv;
import dev.braintrust.instrumentation.ResponseToolSpans;
import dev.braintrust.json.BraintrustJsonMapper;
import io.opentelemetry.api.OpenTelemetry;
import io.opentelemetry.api.trace.Span;
Expand Down Expand Up @@ -130,7 +131,7 @@ public void close() {
var response = underlying.execute(bufferedRequest, requestOptions);
// Always tee the response body. onStreamClosed() detects whether the collected
// bytes are SSE or plain JSON and tags the span accordingly.
return new TeeingStreamHttpResponse(response, span);
return new TeeingStreamHttpResponse(response, span, tracer);
} catch (Exception e) {
InstrumentationSemConv.tagLLMSpanResponse(span, e);
span.end();
Expand Down Expand Up @@ -159,7 +160,9 @@ public void close() {
return underlying
.executeAsync(bufferedRequest, requestOptions)
.thenApply(
response -> (HttpResponse) new TeeingStreamHttpResponse(response, span))
response ->
(HttpResponse)
new TeeingStreamHttpResponse(response, span, tracer))
.whenComplete(
(response, t) -> {
if (t != null) {
Expand Down Expand Up @@ -239,21 +242,22 @@ private static String readBodyAsString(HttpRequestBody body) {
* the bytes are an SSE stream (first non-empty line starts with {@code "data: "}) or a plain
* JSON response, and parses accordingly.
*/
private static void tagSpanFromBuffer(Span span, byte[] bytes, Long timeToFirstTokenNanos) {
if (bytes.length == 0) return;
private static String tagSpanFromBuffer(Span span, byte[] bytes, Long timeToFirstTokenNanos) {
if (bytes.length == 0) return null;
try {
String firstLine = firstNonEmptyLine(bytes);
if (firstLine != null
&& (firstLine.startsWith("data:") || firstLine.startsWith("event:"))) {
tagSpanFromSseBytes(span, bytes, timeToFirstTokenNanos);
return tagSpanFromSseBytes(span, bytes, timeToFirstTokenNanos);
} else {
String responseJson = new String(bytes, StandardCharsets.UTF_8);
InstrumentationSemConv.tagLLMSpanResponse(
span,
InstrumentationSemConv.PROVIDER_NAME_OPENAI,
new String(bytes, StandardCharsets.UTF_8));
span, InstrumentationSemConv.PROVIDER_NAME_OPENAI, responseJson);
return responseJson;
}
} catch (Exception e) {
log.error("Could not tag span from response buffer", e);
return null;
}
}

Expand All @@ -273,7 +277,7 @@ private static String firstNonEmptyLine(byte[] bytes) {
* Parses SSE wire bytes, feeds each {@code data:} chunk through {@link
* ChatCompletionAccumulator}, then tags the span with the reassembled output JSON.
*/
private static void tagSpanFromSseBytes(
private static String tagSpanFromSseBytes(
Span span, byte[] sseBytes, Long timeToFirstTokenNanos) {
try {
var reader =
Expand Down Expand Up @@ -330,8 +334,10 @@ private static void tagSpanFromSseBytes(
responseJson,
timeToFirstTokenNanos);
}
return responseJson;
} catch (Exception e) {
log.error("Could not parse SSE buffer to tag streaming span output", e);
return null;
}
}

Expand All @@ -343,14 +349,16 @@ private static void tagSpanFromSseBytes(
private static final class TeeingStreamHttpResponse implements HttpResponse {
private final HttpResponse delegate;
private final Span span;
private final Tracer tracer;
private final long spanStartNanos = System.nanoTime();
private final AtomicLong timeToFirstTokenNanos = new AtomicLong();
private final ByteArrayOutputStream teeBuffer = new ByteArrayOutputStream();
private final InputStream teeStream;

TeeingStreamHttpResponse(HttpResponse delegate, Span span) {
TeeingStreamHttpResponse(HttpResponse delegate, Span span, Tracer tracer) {
this.delegate = delegate;
this.span = span;
this.tracer = tracer;
this.teeStream =
new TeeInputStream(
delegate.body(), teeBuffer, this::onFirstByte, this::onStreamClosed);
Expand All @@ -369,7 +377,13 @@ private void onStreamClosed() {
synchronized (teeBuffer) {
bytes = teeBuffer.toByteArray();
}
tagSpanFromBuffer(span, bytes, timeToFirstTokenNanos.get());
String responseJson = tagSpanFromBuffer(span, bytes, timeToFirstTokenNanos.get());
if (responseJson != null) {
// Emit child tool spans (web search, etc.) parented to the LLM span, while it
// is still live. No-op for Chat Completions responses (no `output` array).
ResponseToolSpans.emitOpenAIResponseToolSpans(
tracer, Context.current().with(span), responseJson);
}
} finally {
span.end();
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -545,10 +545,18 @@ void testResponsesStreamingWithTools() {
.toList();
assertFalse(functionCalls.isEmpty(), "model should call a function tool");

var spans = testHarness.awaitExportedSpans();
assertEquals(1, spans.size());
var span = spans.get(0);
// The LLM span plus one child tool span per function_call in the response output.
var spans = testHarness.awaitExportedSpans(1 + functionCalls.size());
var span = llmSpan(spans);
assertValidOpenAISpan(span, true);
assertEquals(
functionCalls.size(),
toolSpans(spans).size(),
"each function_call should produce a child tool span");
assertTrue(
toolSpans(spans).stream()
.allMatch(t -> t.getParentSpanId().equals(span.getSpanId())),
"tool spans must be children of the LLM span");

// The Responses API serializes output as an "output" array; the streamed function_call
// items must survive reconstruction into the span exactly as the client accumulated them.
Expand Down Expand Up @@ -609,6 +617,26 @@ private static FunctionTool.Parameters weatherParameters(
.build();
}

/** The single LLM span among exported spans (span_attributes type == "llm"). */
@SneakyThrows
private static SpanData llmSpan(List<SpanData> spans) {
var llm = spans.stream().filter(s -> isSpanType(s, "llm")).toList();
assertEquals(1, llm.size(), "expected exactly one LLM span");
return llm.get(0);
}

/** Child tool spans (span_attributes type == "tool") among exported spans. */
private static List<SpanData> toolSpans(List<SpanData> spans) {
return spans.stream().filter(s -> isSpanType(s, "tool")).toList();
}

@SneakyThrows
private static boolean isSpanType(SpanData span, String type) {
String attr =
span.getAttributes().get(AttributeKey.stringKey("braintrust.span_attributes"));
return attr != null && type.equals(JSON_MAPPER.readTree(attr).path("type").asText());
}

@SneakyThrows
private JsonNode spanOutput(SpanData span) {
String outputJson =
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,143 @@
package dev.braintrust.instrumentation.openai.v2_15_0;

import static org.junit.jupiter.api.Assertions.*;

import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.openai.client.OpenAIClient;
import com.openai.client.okhttp.OpenAIOkHttpClient;
import com.openai.core.http.StreamResponse;
import com.openai.helpers.ResponseAccumulator;
import com.openai.models.ChatModel;
import com.openai.models.responses.EasyInputMessage;
import com.openai.models.responses.Response;
import com.openai.models.responses.ResponseCreateParams;
import com.openai.models.responses.ResponseInputItem;
import com.openai.models.responses.ResponseStreamEvent;
import com.openai.models.responses.WebSearchTool;
import dev.braintrust.TestHarness;
import dev.braintrust.instrumentation.Instrumenter;
import io.opentelemetry.api.common.AttributeKey;
import io.opentelemetry.sdk.trace.data.SpanData;
import java.util.List;
import lombok.SneakyThrows;
import net.bytebuddy.agent.ByteBuddyAgent;
import org.junit.jupiter.api.BeforeAll;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;

/**
* Verifies that built-in web-search tool calls in the OpenAI Responses API are captured as child
* {@code type:"tool"} spans parented to the LLM span, giving web search its own cost/latency
* visibility on the trace timeline.
*/
public class BraintrustOpenAIWebSearchTest {
private static final ObjectMapper JSON_MAPPER = new ObjectMapper();
private static final AttributeKey<String> SPAN_ATTRIBUTES =
AttributeKey.stringKey("braintrust.span_attributes");
private static final AttributeKey<String> METADATA =
AttributeKey.stringKey("braintrust.metadata");

@BeforeAll
public static void beforeAll() {
var instrumentation = ByteBuddyAgent.install();
Instrumenter.install(instrumentation, BraintrustOpenAIWebSearchTest.class.getClassLoader());
}

private TestHarness testHarness;

@BeforeEach
void beforeEach() {
testHarness = TestHarness.setup();
}

private static ResponseCreateParams webSearchRequest() {
return ResponseCreateParams.builder()
.model(ChatModel.GPT_4O)
.inputOfResponse(
List.of(
ResponseInputItem.ofEasyInputMessage(
EasyInputMessage.builder()
.role(EasyInputMessage.Role.USER)
.content(
"What is one recent headline about"
+ " artificial intelligence? Use web"
+ " search.")
.build())))
.addTool(
WebSearchTool.builder().type(WebSearchTool.Type.WEB_SEARCH_PREVIEW).build())
.build();
}

@Test
@SneakyThrows
void testResponsesWebSearch() {
OpenAIClient client =
OpenAIOkHttpClient.builder()
.baseUrl(testHarness.openAiBaseUrl())
.apiKey(testHarness.openAiApiKey())
.build();

Response response = client.responses().create(webSearchRequest());
assertNotNull(response);

var spans = testHarness.awaitExportedSpans(2);
assertWebSearchToolSpans(spans);
}

@Test
@SneakyThrows
void testResponsesWebSearchStreaming() {
OpenAIClient client =
OpenAIOkHttpClient.builder()
.baseUrl(testHarness.openAiBaseUrl())
.apiKey(testHarness.openAiApiKey())
.build();

var accumulator = ResponseAccumulator.create();
try (StreamResponse<ResponseStreamEvent> stream =
client.responses().createStreaming(webSearchRequest())) {
stream.stream().forEach(accumulator::accumulate);
}
assertFalse(accumulator.response().output().isEmpty(), "should generate a response");

var spans = testHarness.awaitExportedSpans(2);
assertWebSearchToolSpans(spans);
}

@SneakyThrows
private static void assertWebSearchToolSpans(List<SpanData> spans) {
// Exactly one LLM span (the Responses request), plus one or more tool spans.
var llmSpans = spans.stream().filter(s -> isType(s, "llm")).toList();
assertEquals(1, llmSpans.size(), "expected a single LLM span");
var llm = llmSpans.get(0);

var webSearchSpans =
spans.stream()
.filter(s -> isType(s, "tool"))
.filter(s -> "web_search_call".equals(s.getName()))
.toList();
assertFalse(
webSearchSpans.isEmpty(),
"expected at least one web_search_call tool span, got spans: "
+ spans.stream().map(SpanData::getName).toList());

for (var ws : webSearchSpans) {
assertEquals(
llm.getSpanId(),
ws.getParentSpanId(),
"web_search_call tool span must be a child of the LLM span");
JsonNode metadata = JSON_MAPPER.readTree(ws.getAttributes().get(METADATA));
assertEquals("web_search_call", metadata.path("tool_type").asText());
}
}

@SneakyThrows
private static boolean isType(SpanData span, String type) {
String attr = span.getAttributes().get(SPAN_ATTRIBUTES);
if (attr == null) {
return false;
}
return type.equals(JSON_MAPPER.readTree(attr).path("type").asText());
}
}
Loading
Loading