Skip to content

Commit b782a54

Browse files
committed
feat(tracing): emit real token usage from openai_agents temporal models
Captures ResponseCompletedEvent usage in the streaming model (was zeroed) and response.usage in both tracing wrappers, writing span.output.usage for billing. Also implements stream_response on the tracing wrappers, which were abstract and raised TypeError on instantiation.
1 parent aa51ab4 commit b782a54

6 files changed

Lines changed: 294 additions & 12 deletions

File tree

src/agentex/lib/core/temporal/plugins/openai_agents/models/temporal_streaming_model.py

Lines changed: 30 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,7 @@
5454

5555
# AgentEx SDK imports
5656
from agentex.lib import adk
57+
from agentex.lib.core.tracing.usage import usage_from_openai_response_usage
5758
from agentex.lib.core.tracing.tracer import AsyncTracer
5859
from agentex.types.task_message_delta import TextDelta, ReasoningContentDelta, ReasoningSummaryDelta
5960
from agentex.types.task_message_update import StreamTaskMessageFull, StreamTaskMessageDelta
@@ -601,6 +602,7 @@ async def get_response(
601602
reasoning_summaries = []
602603
reasoning_contents = []
603604
event_count = 0
605+
response_usage = None
604606

605607
# We expect task_id to always be provided for streaming
606608
if not task_id:
@@ -806,6 +808,8 @@ async def get_response(
806808
# Use the final output from the response
807809
output_items = response.output
808810
logger.debug(f"[TemporalStreamingModel] Found {len(output_items)} output items in final response")
811+
if response is not None:
812+
response_usage = getattr(response, 'usage', None)
809813

810814
# End of event processing loop - close any open contexts
811815
if reasoning_context:
@@ -844,14 +848,27 @@ async def get_response(
844848
)
845849
response_output.append(message)
846850

847-
# Create usage object
848-
usage = Usage(
849-
input_tokens=0,
850-
output_tokens=0,
851-
total_tokens=0,
852-
input_tokens_details=InputTokensDetails(cached_tokens=0),
853-
output_tokens_details=OutputTokensDetails(reasoning_tokens=len(''.join(reasoning_contents)) // 4), # Approximate
854-
)
851+
# Create usage object from the final response's real usage
852+
if response_usage is not None:
853+
usage = Usage(
854+
requests=1,
855+
input_tokens=response_usage.input_tokens or 0,
856+
output_tokens=response_usage.output_tokens or 0,
857+
total_tokens=response_usage.total_tokens or 0,
858+
input_tokens_details=response_usage.input_tokens_details
859+
or InputTokensDetails(cached_tokens=0),
860+
output_tokens_details=response_usage.output_tokens_details
861+
or OutputTokensDetails(reasoning_tokens=0),
862+
)
863+
else:
864+
# No usage reported by the API (e.g. stream ended early)
865+
usage = Usage(
866+
input_tokens=0,
867+
output_tokens=0,
868+
total_tokens=0,
869+
input_tokens_details=InputTokensDetails(cached_tokens=0),
870+
output_tokens_details=OutputTokensDetails(reasoning_tokens=len(''.join(reasoning_contents)) // 4), # Approximate
871+
)
855872

856873
# Serialize response output items for span tracing
857874
new_items = []
@@ -907,7 +924,11 @@ async def get_response(
907924
# Include tool outputs if any were processed
908925
if tool_outputs:
909926
output_data["tool_outputs"] = tool_outputs
910-
927+
# Per-call usage for billing; deduped against any turn aggregate
928+
usage_blob = usage_from_openai_response_usage(response_usage)
929+
if usage_blob:
930+
output_data["usage"] = usage_blob
931+
911932
span.output = output_data
912933

913934
# Return the response

src/agentex/lib/core/temporal/plugins/openai_agents/models/temporal_tracing_model.py

Lines changed: 26 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
from agents.models.openai_responses import OpenAIResponsesModel
2929
from agents.models.openai_chatcompletions import OpenAIChatCompletionsModel
3030

31+
from agentex.lib.core.tracing.usage import usage_from_openai_response_usage
3132
from agentex.lib.core.tracing.tracer import AsyncTracer
3233

3334
# Import AgentEx components
@@ -243,10 +244,15 @@ async def get_response(
243244
continue
244245

245246
# Set span output with structured data
246-
span.output = { # type: ignore[attr-defined]
247+
output_data: dict[str, Any] = {
247248
"new_items": new_items,
248249
"final_output": final_output,
249250
}
251+
# Per-call usage for billing; deduped against any turn aggregate
252+
usage_blob = usage_from_openai_response_usage(getattr(response, "usage", None))
253+
if usage_blob:
254+
output_data["usage"] = usage_blob
255+
span.output = output_data # type: ignore[attr-defined]
250256

251257
return response
252258

@@ -271,6 +277,12 @@ async def get_response(
271277
**kwargs,
272278
)
273279

280+
@override
281+
def stream_response(self, *args, **kwargs):
282+
"""Streaming is handled via get_response in Temporal activities.
283+
Required so the class is concrete and instantiable at runtime."""
284+
raise NotImplementedError("stream_response is not used in Temporal activities - use get_response instead")
285+
274286

275287
class TemporalTracingChatCompletionsModel(Model):
276288
"""Wrapper for OpenAIChatCompletionsModel that adds AgentEx tracing.
@@ -376,10 +388,15 @@ async def get_response(
376388
continue
377389

378390
# Set span output with structured data
379-
span.output = { # type: ignore[attr-defined]
391+
output_data: dict[str, Any] = {
380392
"new_items": new_items,
381393
"final_output": final_output,
382394
}
395+
# Per-call usage for billing; deduped against any turn aggregate
396+
usage_blob = usage_from_openai_response_usage(getattr(response, "usage", None))
397+
if usage_blob:
398+
output_data["usage"] = usage_blob
399+
span.output = output_data # type: ignore[attr-defined]
383400

384401
return response
385402

@@ -399,4 +416,10 @@ async def get_response(
399416
handoffs=handoffs,
400417
tracing=tracing,
401418
**kwargs,
402-
)
419+
)
420+
421+
@override
422+
def stream_response(self, *args, **kwargs):
423+
"""Streaming is handled via get_response in Temporal activities.
424+
Required so the class is concrete and instantiable at runtime."""
425+
raise NotImplementedError("stream_response is not used in Temporal activities - use get_response instead")

tests/lib/core/temporal/__init__.py

Whitespace-only changes.

tests/lib/core/temporal/plugins/__init__.py

Whitespace-only changes.

tests/lib/core/temporal/plugins/openai_agents/__init__.py

Whitespace-only changes.
Lines changed: 238 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,238 @@
1+
"""Tests that the openai_agents temporal models copy real token usage onto spans.
2+
3+
The backend bills per-call usage from ``span.output["usage"]``; these tests
4+
assert each model writes the framework-reported usage there (and, for the
5+
streaming model, into the returned ``ModelResponse.usage``) instead of
6+
dropping it.
7+
"""
8+
9+
from __future__ import annotations
10+
11+
from datetime import UTC, datetime
12+
from contextlib import asynccontextmanager
13+
from unittest.mock import AsyncMock, MagicMock, patch
14+
15+
import pytest
16+
from agents import ModelResponse
17+
from agents.usage import Usage, InputTokensDetails, OutputTokensDetails
18+
19+
from agentex.types.span import Span
20+
21+
pytestmark = pytest.mark.asyncio
22+
23+
24+
class FakeTrace:
25+
"""Captures spans handed out by trace.span() so tests can inspect them."""
26+
27+
def __init__(self) -> None:
28+
self.spans: list[Span] = []
29+
30+
@asynccontextmanager
31+
async def span(self, name, parent_id=None, input=None, data=None, task_id=None):
32+
span = Span(
33+
id=f"span-{len(self.spans)}",
34+
name=name,
35+
start_time=datetime.now(UTC),
36+
trace_id="trace-1",
37+
parent_id=parent_id,
38+
input=input,
39+
data=data,
40+
task_id=task_id,
41+
)
42+
self.spans.append(span)
43+
yield span
44+
45+
46+
class FakeTracer:
47+
def __init__(self) -> None:
48+
self.trace_obj = FakeTrace()
49+
50+
def trace(self, trace_id):
51+
return self.trace_obj
52+
53+
54+
@pytest.fixture
55+
def tracing_contextvars():
56+
from agentex.lib.core.temporal.plugins.openai_agents.interceptors.context_interceptor import (
57+
streaming_task_id,
58+
streaming_trace_id,
59+
streaming_parent_span_id,
60+
)
61+
62+
tokens = [
63+
streaming_task_id.set("task-1"),
64+
streaming_trace_id.set("trace-1"),
65+
streaming_parent_span_id.set("parent-span-1"),
66+
]
67+
yield
68+
streaming_task_id.reset(tokens[0])
69+
streaming_trace_id.reset(tokens[1])
70+
streaming_parent_span_id.reset(tokens[2])
71+
72+
73+
def _agents_usage() -> Usage:
74+
return Usage(
75+
requests=1,
76+
input_tokens=120,
77+
output_tokens=80,
78+
total_tokens=200,
79+
input_tokens_details=InputTokensDetails(cached_tokens=30),
80+
output_tokens_details=OutputTokensDetails(reasoning_tokens=40),
81+
)
82+
83+
EXPECTED_USAGE_BLOB = {
84+
"input_tokens": 120,
85+
"output_tokens": 80,
86+
"total_tokens": 200,
87+
"cached_input_tokens": 30,
88+
"reasoning_tokens": 40,
89+
}
90+
91+
92+
class TestTemporalTracingModels:
93+
async def _run_wrapper(self, wrapper_cls) -> Span:
94+
tracer = FakeTracer()
95+
base_model = MagicMock()
96+
base_model.model = "gpt-4o"
97+
base_model.get_response = AsyncMock(
98+
return_value=ModelResponse(output=[], usage=_agents_usage(), response_id="resp-1")
99+
)
100+
101+
model = wrapper_cls(base_model, tracer)
102+
from agents import ModelSettings
103+
104+
response = await model.get_response(
105+
system_instructions=None,
106+
input="hello",
107+
model_settings=ModelSettings(),
108+
tools=[],
109+
output_schema=None,
110+
handoffs=[],
111+
tracing=None,
112+
)
113+
assert response.usage.input_tokens == 120
114+
assert len(tracer.trace_obj.spans) == 1
115+
return tracer.trace_obj.spans[0]
116+
117+
async def test_responses_model_writes_usage_to_span_output(self, tracing_contextvars):
118+
from agentex.lib.core.temporal.plugins.openai_agents.models.temporal_tracing_model import (
119+
TemporalTracingResponsesModel,
120+
)
121+
122+
span = await self._run_wrapper(TemporalTracingResponsesModel)
123+
assert span.output["usage"] == EXPECTED_USAGE_BLOB
124+
125+
async def test_chat_completions_model_writes_usage_to_span_output(self, tracing_contextvars):
126+
from agentex.lib.core.temporal.plugins.openai_agents.models.temporal_tracing_model import (
127+
TemporalTracingChatCompletionsModel,
128+
)
129+
130+
span = await self._run_wrapper(TemporalTracingChatCompletionsModel)
131+
assert span.output["usage"] == EXPECTED_USAGE_BLOB
132+
133+
134+
class FakeStream:
135+
def __init__(self, events) -> None:
136+
self._events = events
137+
138+
def __aiter__(self):
139+
async def gen():
140+
for event in self._events:
141+
yield event
142+
143+
return gen()
144+
145+
146+
class TestTemporalStreamingModel:
147+
async def test_streaming_model_captures_final_response_usage(self, tracing_contextvars):
148+
import agentex.lib.core.temporal.plugins.openai_agents.models.temporal_streaming_model as tsm
149+
150+
from openai.types.responses import Response, ResponseCompletedEvent
151+
from openai.types.responses.response_usage import (
152+
ResponseUsage,
153+
InputTokensDetails as ResponseInputTokensDetails,
154+
OutputTokensDetails as ResponseOutputTokensDetails,
155+
)
156+
157+
usage = ResponseUsage(
158+
input_tokens=120,
159+
output_tokens=80,
160+
total_tokens=200,
161+
input_tokens_details=ResponseInputTokensDetails(cached_tokens=30),
162+
output_tokens_details=ResponseOutputTokensDetails(reasoning_tokens=40),
163+
)
164+
completed = ResponseCompletedEvent.model_construct(
165+
type="response.completed",
166+
response=Response.model_construct(output=[], usage=usage),
167+
)
168+
169+
fake_tracer = FakeTracer()
170+
openai_client = MagicMock()
171+
openai_client.responses.create = AsyncMock(return_value=FakeStream([completed]))
172+
173+
with (
174+
patch.object(tsm, "create_async_agentex_client", return_value=MagicMock()),
175+
patch.object(tsm, "AsyncTracer", return_value=fake_tracer),
176+
):
177+
model = tsm.TemporalStreamingModel(model_name="gpt-4o", openai_client=openai_client)
178+
179+
from agents import ModelSettings
180+
181+
response = await model.get_response(
182+
system_instructions=None,
183+
input="hello",
184+
model_settings=ModelSettings(),
185+
tools=[],
186+
output_schema=None,
187+
handoffs=[],
188+
tracing=None,
189+
)
190+
191+
# Real usage lands on the returned ModelResponse (was zeroed before)
192+
assert response.usage.requests == 1
193+
assert response.usage.input_tokens == 120
194+
assert response.usage.output_tokens == 80
195+
assert response.usage.total_tokens == 200
196+
assert response.usage.input_tokens_details.cached_tokens == 30
197+
assert response.usage.output_tokens_details.reasoning_tokens == 40
198+
199+
# And on the span output for billing
200+
assert len(fake_tracer.trace_obj.spans) == 1
201+
span = fake_tracer.trace_obj.spans[0]
202+
assert span.output["usage"] == EXPECTED_USAGE_BLOB
203+
204+
async def test_streaming_model_omits_usage_when_api_reports_none(self, tracing_contextvars):
205+
import agentex.lib.core.temporal.plugins.openai_agents.models.temporal_streaming_model as tsm
206+
207+
from openai.types.responses import Response, ResponseCompletedEvent
208+
209+
completed = ResponseCompletedEvent.model_construct(
210+
type="response.completed",
211+
response=Response.model_construct(output=[], usage=None),
212+
)
213+
214+
fake_tracer = FakeTracer()
215+
openai_client = MagicMock()
216+
openai_client.responses.create = AsyncMock(return_value=FakeStream([completed]))
217+
218+
with (
219+
patch.object(tsm, "create_async_agentex_client", return_value=MagicMock()),
220+
patch.object(tsm, "AsyncTracer", return_value=fake_tracer),
221+
):
222+
model = tsm.TemporalStreamingModel(model_name="gpt-4o", openai_client=openai_client)
223+
224+
from agents import ModelSettings
225+
226+
response = await model.get_response(
227+
system_instructions=None,
228+
input="hello",
229+
model_settings=ModelSettings(),
230+
tools=[],
231+
output_schema=None,
232+
handoffs=[],
233+
tracing=None,
234+
)
235+
236+
assert response.usage.input_tokens == 0
237+
span = fake_tracer.trace_obj.spans[0]
238+
assert "usage" not in span.output

0 commit comments

Comments
 (0)