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
37 changes: 29 additions & 8 deletions py/src/braintrust/integrations/openai/test_openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -448,18 +448,32 @@ async def collect_events():
stream = await parse_result
assert stream.response
async with stream:
return [event async for event in stream]
if Version(openai.__version__) >= Version("3.23.0"):
result = await stream.get_final_result()
return None, result
return [event async for event in stream], None

events = asyncio.run(collect_events())
events, result = asyncio.run(collect_events())
else:
raw_response = sessions.with_raw_response.create(**params)
assert raw_response.headers
with raw_response.parse() as stream:
assert stream.response
events = list(stream)

completed = next(
event for event in events if event.type == "agent.session.turn.completed" and event.turn.subagent_id is None
if Version(openai.__version__) >= Version("3.23.0"):
result = stream.get_final_result()
events = None
else:
events = list(stream)
result = None

completed = (
next(
event
for event in events
if event.type == "agent.session.turn.completed" and event.turn.subagent_id is None
)
if events is not None
else None
)

spans = memory_logger.pop()
Expand All @@ -471,11 +485,18 @@ async def collect_events():
assert task_span["span_attributes"]["name"] == "openai.agents.sessions.create"
assert task_span["input"] == input_text
assert expected_text in task_span["output"]
if result is not None:
assert expected_text in result.output_text
expected_session_id = result.session_id
expected_turn_id = result.turn_id
else:
expected_session_id = completed.session_id
expected_turn_id = completed.turn_id
assert task_span["metadata"]["provider"] == "openai"
assert task_span["metadata"]["model"] == "gpt-6-astra"
assert task_span["metadata"]["environment_type"] == environment_type
assert task_span["metadata"]["session_id"] == completed.session_id
assert task_span["metadata"]["turn_id"] == completed.turn_id
assert task_span["metadata"]["session_id"] == expected_session_id
assert task_span["metadata"]["turn_id"] == expected_turn_id
assert task_span["metadata"]["status"] == "completed"
assert task_span["context"]["span_origin"]["instrumentation"]["name"] == "openai-auto"
assert task_span["metrics"]["time_to_first_token"] >= 0
Expand Down
35 changes: 33 additions & 2 deletions py/src/braintrust/integrations/openai/tracing.py
Original file line number Diff line number Diff line change
Expand Up @@ -996,6 +996,37 @@ async def __aexit__(self, exc_type, exc_val, exc_tb) -> Any:
return result


class _TracedAgentSessionStream(_TracedStream):
"""Agent session stream methods must drain through the traced iterator."""

def with_result_collection(self) -> Any:
self._wrapped.with_result_collection()
return self

def get_final_result(self) -> Any:
self._wrapped.with_result_collection()
for _ in self:
pass
return self._wrapped.get_final_result()
Comment thread
AbhiPrasad marked this conversation as resolved.


class _AsyncTracedAgentSessionStream(_AsyncTracedStream):
"""Async agent session stream methods must drain through the traced iterator."""

def with_result_collection(self) -> Any:
self._wrapped.with_result_collection()
return self

async def get_final_result(self) -> Any:
self._wrapped.with_result_collection()
async for _ in self:
pass
result = self._wrapped.get_final_result()
if inspect.isawaitable(result):
return await result
return result


class _RawResponseWithTracedStream(NamedWrapper):
"""Proxy for LegacyAPIResponse that replaces parse() with a traced stream,
so that with_raw_response + stream=True preserves both headers and tracing."""
Expand Down Expand Up @@ -1420,7 +1451,7 @@ def gen():
finally:
trace.finish()

traced_stream = _TracedStream(stream, gen(), trace.finish)
traced_stream = _TracedAgentSessionStream(stream, gen(), trace.finish)
if raw_requested:
return _RawResponseWithTracedStream(create_response, traced_stream)
return traced_stream
Expand Down Expand Up @@ -1463,7 +1494,7 @@ async def gen():
finally:
trace.finish()

traced_stream = _AsyncTracedStream(stream, gen(), trace.finish)
traced_stream = _AsyncTracedAgentSessionStream(stream, gen(), trace.finish)
if raw_requested:
return _RawResponseWithTracedStream(
create_response,
Expand Down
Loading