fix(core): count Responses requests without usage (#4453)

Co-authored-by: ayaangazali <ayaangazali.work@gmail.com>
This commit is contained in:
Kazuhiro Sera
2026-08-16 10:28:43 +09:00
committed by GitHub
parent cb8a2e7e7d
commit 94da8ed49c
5 changed files with 542 additions and 60 deletions
@@ -39,6 +39,7 @@ from ....models._run_context import get_model_run_owner
from ....models.openai_responses import OpenAIResponsesModel, _is_openai_omitted_value
from ....tool import Tool
from ....tool_context import ToolContext
from ....usage import _mark_requests_completed_without_usage
_BETA_ID = "responses_multi_agent=v1"
_ROOT_AGENT_NAME = "/root"
@@ -784,6 +785,11 @@ class OpenAIHostedMultiAgentModel(OpenAIResponsesModel):
request_usages=active.request_usages,
request_count=active.request_count,
)
if normalized_response.usage is None:
_mark_requests_completed_without_usage(
normalized_response,
active.request_count,
)
payload = _model_dump(completed_event)
payload["response"] = normalized_response
normalized_event = _construct_event("response.completed", payload)
+90 -24
View File
@@ -5,7 +5,7 @@ import contextlib
import inspect
import json
import weakref
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping, Sequence
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping, Sequence
from contextvars import ContextVar
from dataclasses import asdict, dataclass, is_dataclass
from enum import Enum
@@ -78,7 +78,9 @@ from ..tracing import SpanError, response_span
from ..usage import (
Usage,
_attach_raw_usage_snapshot,
_mark_request_completed_without_usage,
_raw_usage_snapshot,
_requests_for_response_without_usage,
_response_usage_to_usage,
model_usage_to_span_usage,
)
@@ -229,6 +231,23 @@ class OpenAIResponsesWebSocketOptions(TypedDict):
"""
def _mark_transport_request_without_usage(response: object) -> None:
"""Mark one adapter-owned request when its completed response omits usage."""
if isinstance(response, Response) and response.usage is None:
_mark_request_completed_without_usage(response)
def _usage_from_response(response: Response) -> Usage:
"""Convert provider usage while preserving an adapter-owned request marker."""
if response.usage is not None:
return _response_usage_to_usage(response.usage)
return Usage(requests=_requests_for_response_without_usage(response))
async def _no_stream_cleanup() -> None:
"""Provide a no-op cleanup callback for an externally owned stream."""
class _ResponseStreamWithRequestId:
"""Wrap an SDK event stream and retain the originating request ID."""
@@ -241,18 +260,20 @@ class _ResponseStreamWithRequestId:
def __init__(
self,
stream: AsyncIterator[ResponseStreamEvent],
stream: AsyncIterator[ResponseStreamEvent] | AsyncIterable[ResponseStreamEvent],
*,
request_id: str | None,
cleanup: Callable[[], Awaitable[object]],
) -> None:
self._stream = stream
self._source = stream
self._stream: AsyncIterator[ResponseStreamEvent] | None = None
self.request_id = request_id
self._cleanup = cleanup
self._closed = False
self._stream_close_complete = False
self._cleanup_complete = False
self._yielded_terminal_event = False
self._close_task: asyncio.Future[None] | None = None
def __aiter__(self) -> _ResponseStreamWithRequestId:
return self
@@ -261,8 +282,13 @@ class _ResponseStreamWithRequestId:
if self._closed:
raise StopAsyncIteration
stream = self._stream
if stream is None:
stream = self._source.__aiter__()
self._stream = stream
try:
event = await self._stream.__anext__()
event = await stream.__anext__()
except StopAsyncIteration:
self._closed = True
await self._cleanup_after_exhaustion()
@@ -272,14 +298,13 @@ class _ResponseStreamWithRequestId:
event_type = getattr(event, "type", None)
if event_type in self._TERMINAL_EVENT_TYPES:
self._yielded_terminal_event = True
if event_type == "response.completed":
_mark_transport_request_without_usage(getattr(event, "response", None))
return event
async def aclose(self) -> None:
self._closed = True
try:
await self._close_stream_once()
finally:
await self._cleanup_once()
await self._close_stream_and_cleanup()
async def close(self) -> None:
await self.aclose()
@@ -305,7 +330,7 @@ class _ResponseStreamWithRequestId:
async def _cleanup_after_exhaustion(self) -> None:
try:
await self._cleanup_once()
await self._close_stream_and_cleanup()
except Exception as exc:
if self._yielded_terminal_event:
log_model_action_debug(
@@ -314,17 +339,54 @@ class _ResponseStreamWithRequestId:
return
raise
async def _close_stream_and_cleanup(self) -> None:
if self._close_task is None:
self._close_task = asyncio.ensure_future(self._finish_stream_close_and_cleanup())
await asyncio.shield(self._close_task)
async def _finish_stream_close_and_cleanup(self) -> None:
try:
await self._close_stream_once()
except BaseException:
with contextlib.suppress(BaseException):
await self._cleanup_once()
raise
await self._cleanup_once()
async def _close_stream_once(self) -> None:
if self._stream_close_complete:
return
self._stream_close_complete = True
aclose = getattr(self._stream, "aclose", None)
# An async iterable may create a separate iterator that owns its cleanup. Close both the
# resolved iterator and the source object, while avoiding duplicate close calls when they
# are the same object.
resolved = self._stream
if self._source is resolved:
await self._close_object(resolved)
return
try:
await self._close_object(resolved)
except BaseException:
# The source may own transport cleanup independently of its iterator. Preserve the
# iterator's original error while still making a best effort to release the source.
with contextlib.suppress(BaseException):
await self._close_object(self._source)
raise
await self._close_object(self._source)
@staticmethod
async def _close_object(target: object | None) -> None:
if target is None:
return
aclose = getattr(target, "aclose", None)
if callable(aclose):
await aclose()
return
close = getattr(self._stream, "close", None)
close = getattr(target, "close", None)
if callable(close):
close_result = close()
if inspect.isawaitable(close_result):
@@ -528,12 +590,8 @@ class OpenAIResponsesModel(Model):
),
)
usage = (
_response_usage_to_usage(response.usage)
if response.usage is not None
else Usage()
)
if response.usage is not None:
usage = _usage_from_response(response)
if response.usage is not None or usage.requests:
span_response.span_data.usage = model_usage_to_span_usage(usage)
if tracing.include_data():
@@ -610,6 +668,11 @@ class OpenAIResponsesModel(Model):
final_response = chunk.response
if model_settings.preserve_raw_usage is True:
_attach_raw_usage_snapshot(chunk.response, chunk.response.usage)
usage = _usage_from_response(chunk.response)
if chunk.response.usage is not None or usage.requests:
# Record before yielding the terminal event because consumers may
# close the generator immediately after receiving it.
span_response.span_data.usage = model_usage_to_span_usage(usage)
elif chunk_type in {
"response.failed",
"response.incomplete",
@@ -669,11 +732,6 @@ class OpenAIResponsesModel(Model):
if final_response is not None and tracing.include_data():
span_response.span_data.response = final_response
span_response.span_data.input = input
if final_response is not None and final_response.usage is not None:
span_response.span_data.usage = model_usage_to_span_usage(
_response_usage_to_usage(final_response.usage)
)
except Exception as e:
span_response.set_error(
SpanError(
@@ -747,15 +805,21 @@ class OpenAIResponsesModel(Model):
if not stream:
response = await client.responses.create(**create_kwargs)
_mark_transport_request_without_usage(response)
return cast(Response, response)
streaming_response = getattr(client.responses, "with_streaming_response", None)
stream_create = getattr(streaming_response, "create", None)
if not callable(stream_create):
# Some tests and custom clients only implement `responses.create()`. Fall back to the
# older path in that case and simply omit request IDs for streamed calls.
# older path in that case and simply omit request IDs for streamed calls. Keep it in
# the existing stream wrapper so terminal request accounting stays transport-owned.
response = await client.responses.create(**create_kwargs)
return cast(AsyncIterator[ResponseStreamEvent], response)
return _ResponseStreamWithRequestId(
cast(AsyncIterator[ResponseStreamEvent], response),
request_id=None,
cleanup=_no_stream_cleanup,
)
# Keep the raw API response open while callers consume the SSE stream so we can expose
# its request ID on terminal response payloads before cleanup closes the transport.
@@ -1277,6 +1341,8 @@ class OpenAIResponsesWSModel(OpenAIResponsesModel):
}
if is_terminal_event:
yielded_terminal_event = True
if event_type == "response.completed":
_mark_transport_request_without_usage(getattr(event, "response", None))
yield event
if is_terminal_event:
+14 -4
View File
@@ -321,16 +321,26 @@ def _mark_request_completed_without_usage(response: Any) -> None:
Adapters call this instead of synthesizing a zero-filled usage payload, so the raw
provider usage stays absent while the request itself is still counted.
"""
object.__setattr__(response, _REQUEST_WITHOUT_USAGE_ATTR, True)
_mark_requests_completed_without_usage(response, 1)
def _mark_requests_completed_without_usage(response: Any, requests: int) -> None:
"""Record an adapter-owned physical request count without provider usage."""
if requests < 1:
raise ValueError("Completed request count must be at least one.")
object.__setattr__(response, _REQUEST_WITHOUT_USAGE_ATTR, requests)
def _requests_for_response_without_usage(response: Any) -> int:
"""How many requests a usage-less response represents.
Defaults to zero so adapters that multiplex several provider responses into one
response, and report their counts separately, are not double-counted.
Defaults to zero so adapters must explicitly opt in for both singular responses and
responses that aggregate several physical provider requests.
"""
return 1 if getattr(response, _REQUEST_WITHOUT_USAGE_ATTR, False) else 0
requests = getattr(response, _REQUEST_WITHOUT_USAGE_ATTR, 0)
if requests is True:
return 1
return requests if type(requests) is int and requests > 0 else 0
def _response_usage_to_usage(response_usage: Any) -> Usage:
@@ -510,8 +510,10 @@ async def test_runner_injects_two_subagent_calls_into_one_active_response() -> N
@pytest.mark.asyncio
@pytest.mark.parametrize("streamed", [False, True])
@pytest.mark.parametrize("reports_usage", [False, True])
async def test_injection_failure_after_completion_starts_continuation_response(
streamed: bool,
reports_usage: bool,
) -> None:
function_call = {
"id": "fc_fallback",
@@ -528,13 +530,14 @@ async def test_injection_failure_after_completion_starts_continuation_response(
"output": "document:alpha",
}
completed_first_response = _response("resp_old", [function_call])
completed_first_response.usage = _usage(
input_tokens=12,
cached_tokens=3,
cache_write_tokens=1,
output_tokens=4,
reasoning_tokens=2,
)
if reports_usage:
completed_first_response.usage = _usage(
input_tokens=12,
cached_tokens=3,
cache_write_tokens=1,
output_tokens=4,
reasoning_tokens=2,
)
first_events = [
_created("resp_old"),
_done(function_call, sequence_number=2, output_index=0),
@@ -556,13 +559,14 @@ async def test_injection_failure_after_completion_starts_continuation_response(
]
final_message = _root_final_message("continued")
completed_second_response = _response("resp_new", [final_message])
completed_second_response.usage = _usage(
input_tokens=7,
cached_tokens=2,
cache_write_tokens=4,
output_tokens=3,
reasoning_tokens=1,
)
if reports_usage:
completed_second_response.usage = _usage(
input_tokens=7,
cached_tokens=2,
cache_write_tokens=4,
output_tokens=3,
reasoning_tokens=1,
)
second_events = [
_created("resp_new"),
_done(final_message, sequence_number=2, output_index=0),
@@ -600,25 +604,30 @@ async def test_injection_failure_after_completion_starts_continuation_response(
)
assert result.final_output == "continued"
assert result.context_wrapper.usage.input_tokens == 19
assert result.context_wrapper.usage.input_tokens_details.cached_tokens == 5
assert (
getattr(
result.context_wrapper.usage.input_tokens_details,
"cache_write_tokens",
None,
if reports_usage:
assert result.context_wrapper.usage.input_tokens == 19
assert result.context_wrapper.usage.input_tokens_details.cached_tokens == 5
assert (
getattr(
result.context_wrapper.usage.input_tokens_details,
"cache_write_tokens",
None,
)
== 5
)
== 5
)
assert result.context_wrapper.usage.output_tokens == 7
assert result.context_wrapper.usage.output_tokens_details.reasoning_tokens == 3
assert result.context_wrapper.usage.total_tokens == 26
assert result.context_wrapper.usage.output_tokens == 7
assert result.context_wrapper.usage.output_tokens_details.reasoning_tokens == 3
assert result.context_wrapper.usage.total_tokens == 26
assert len(result.context_wrapper.usage.request_usage_entries) == 2
assert [
entry.total_tokens for entry in result.context_wrapper.usage.request_usage_entries
] == [16, 10]
else:
assert result.context_wrapper.usage.input_tokens == 0
assert result.context_wrapper.usage.output_tokens == 0
assert result.context_wrapper.usage.total_tokens == 0
assert result.context_wrapper.usage.request_usage_entries == []
assert result.context_wrapper.usage.requests == 2
assert len(result.context_wrapper.usage.request_usage_entries) == 2
assert [entry.total_tokens for entry in result.context_wrapper.usage.request_usage_entries] == [
16,
10,
]
assert len(client.beta.responses.connections) == 1
connection = client.beta.responses.connections[0]
+392 -1
View File
@@ -8,7 +8,7 @@ from typing import Any, cast
import httpx2
import pytest
from openai import NOT_GIVEN, APIConnectionError, AsyncOpenAI, RateLimitError, omit
from openai.types.responses import ResponseCompletedEvent, ResponseErrorEvent
from openai.types.responses import Response, ResponseCompletedEvent, ResponseErrorEvent
from openai.types.responses.response_create_params import ContextManagement, PromptCacheOptions
from openai.types.responses.response_usage import ResponseUsage
from openai.types.shared.reasoning import Reasoning
@@ -40,6 +40,7 @@ from agents.models.openai_responses import (
OpenAIResponsesModel,
OpenAIResponsesWSModel,
ResponsesWebSocketError,
_ResponseStreamWithRequestId,
_should_retry_pre_event_websocket_disconnect,
)
from agents.retry import ModelRetryAdviceRequest
@@ -4424,3 +4425,393 @@ def test_websocket_get_retry_advice_reports_no_response_started_for_stateful_req
assert advice is not None
assert advice.replay_safety == "unsafe"
assert advice.response_started is False
def _response_without_usage() -> Response:
return Response(
id="resp-no-usage",
created_at=0,
model="fake",
object="response",
output=[],
tool_choice="none",
tools=[],
top_p=None,
parallel_tool_calls=False,
usage=None,
)
def _completed_event_without_usage() -> ResponseCompletedEvent:
return ResponseCompletedEvent(
response=_response_without_usage(),
type="response.completed",
sequence_number=0,
)
def _streaming_client_for(events: list[Any]) -> Any:
class IteratorStream:
def __init__(self) -> None:
self._remaining = list(events)
def __aiter__(self) -> IteratorStream:
return self
async def __anext__(self) -> Any:
if not self._remaining:
raise StopAsyncIteration
return self._remaining.pop(0)
async def close(self) -> None:
return None
stream = IteratorStream()
class APIResponse:
request_id = "req-1"
async def parse(self) -> Any:
return stream
class StreamingContextManager:
async def __aenter__(self) -> APIResponse:
return APIResponse()
async def __aexit__(self, exc_type: Any, exc: Any, tb: Any) -> bool:
return False
class Responses:
with_streaming_response = SimpleNamespace(create=lambda **kwargs: StreamingContextManager())
class Client:
responses = Responses()
base_url = httpx2.URL("https://custom.example.test/v1/")
return Client()
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_request_is_counted_when_responses_provider_omits_usage() -> None:
class Responses:
async def create(self, **kwargs: Any) -> Response:
return _response_without_usage()
class Client:
responses = Responses()
base_url = httpx2.URL("https://custom.example.test/v1/")
model = OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, Client()))
response = await model.get_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(preserve_raw_usage=True),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
)
assert response.usage.requests == 1
assert response.usage.total_tokens == 0
assert response.raw_usage is None
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_streamed_request_is_counted_when_responses_provider_omits_usage() -> None:
completed = _completed_event_without_usage()
client = _streaming_client_for([completed])
agent = Agent(
name="test",
model=OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, client)),
)
result = Runner.run_streamed(agent, "hi")
async for _ in result.stream_events():
pass
assert result.context_wrapper.usage.requests == 1
assert result.context_wrapper.usage.total_tokens == 0
assert completed.response.usage is None
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_fallback_response_stream_counts_request_without_usage() -> None:
completed = _completed_event_without_usage()
class IteratorStream:
def __init__(self) -> None:
self._remaining = [completed]
self.close_calls = 0
def __aiter__(self) -> IteratorStream:
return self
async def __anext__(self) -> ResponseCompletedEvent:
if not self._remaining:
raise StopAsyncIteration
return self._remaining.pop()
async def aclose(self) -> None:
self.close_calls += 1
stream = IteratorStream()
class Responses:
async def create(self, **kwargs: Any) -> IteratorStream:
return stream
class Client:
responses = Responses()
base_url = httpx2.URL("https://custom.example.test/v1/")
model = OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, Client()))
result = Runner.run_streamed(Agent(name="test", model=model), "hi")
async for _ in result.stream_events():
pass
assert result.context_wrapper.usage.requests == 1
assert result.context_wrapper.usage.total_tokens == 0
assert stream.close_calls == 1
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_fallback_response_stream_closes_resolved_iterable_iterator() -> None:
class IterableStream:
def __init__(self) -> None:
self.source_closed = False
self.iterator_closed = False
async def __aiter__(self) -> Any:
try:
yield _completed_event_without_usage()
yield _completed_event_without_usage()
finally:
self.iterator_closed = True
async def aclose(self) -> None:
self.source_closed = True
source = IterableStream()
class Responses:
async def create(self, **kwargs: Any) -> IterableStream:
return source
class Client:
responses = Responses()
base_url = httpx2.URL("https://custom.example.test/v1/")
model = OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, Client()))
stream = model.stream_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
)
async for _ in stream:
break
await stream.aclose()
assert source.iterator_closed
assert source.source_closed
@pytest.mark.asyncio
async def test_response_stream_closes_source_when_resolved_iterator_close_fails() -> None:
class FailingIterator:
def __init__(self) -> None:
self._remaining = [_completed_event_without_usage()]
def __aiter__(self) -> FailingIterator:
return self
async def __anext__(self) -> ResponseCompletedEvent:
if not self._remaining:
raise StopAsyncIteration
return self._remaining.pop()
async def aclose(self) -> None:
raise RuntimeError("iterator close failed")
iterator = FailingIterator()
class IterableSource:
def __init__(self) -> None:
self.closed = False
def __aiter__(self) -> FailingIterator:
return iterator
async def aclose(self) -> None:
self.closed = True
raise RuntimeError("source close failed")
source = IterableSource()
async def cleanup() -> None:
return None
stream = _ResponseStreamWithRequestId(source, request_id=None, cleanup=cleanup)
await stream.__anext__()
with pytest.raises(RuntimeError, match="iterator close failed"):
await stream.aclose()
assert source.closed
@pytest.mark.asyncio
async def test_response_stream_closes_distinct_source_after_normal_exhaustion() -> None:
class IterableSource:
def __init__(self) -> None:
self.source_close_calls = 0
self.iterator_finalized = False
async def __aiter__(self) -> Any:
try:
yield _completed_event_without_usage()
finally:
self.iterator_finalized = True
async def aclose(self) -> None:
self.source_close_calls += 1
source = IterableSource()
cleanup_calls = 0
async def cleanup() -> None:
nonlocal cleanup_calls
cleanup_calls += 1
stream = _ResponseStreamWithRequestId(source, request_id=None, cleanup=cleanup)
async for _ in stream:
pass
assert source.iterator_finalized
assert source.source_close_calls == 1
assert cleanup_calls == 1
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_websocket_request_is_counted_when_responses_provider_omits_usage(
monkeypatch: pytest.MonkeyPatch,
) -> None:
frame = json.dumps(
{
"type": "response.completed",
"response": _response_without_usage().model_dump(),
"sequence_number": 1,
}
)
model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient()))
async def fake_open(
ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None
) -> DummyWSConnection:
return DummyWSConnection([frame])
monkeypatch.setattr(model, "_open_websocket_connection", fake_open)
response = await model.get_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.DISABLED,
)
assert response.usage.requests == 1
assert response.usage.total_tokens == 0
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_websocket_stream_counts_request_without_usage(
monkeypatch: pytest.MonkeyPatch,
) -> None:
frame = json.dumps(
{
"type": "response.completed",
"response": _response_without_usage().model_dump(),
"sequence_number": 1,
}
)
model = OpenAIResponsesWSModel(model="gpt-4", openai_client=cast(Any, DummyWSClient()))
async def fake_open(
ws_url: str, headers: dict[str, str], *, connect_timeout: float | None = None
) -> DummyWSConnection:
return DummyWSConnection([frame])
monkeypatch.setattr(model, "_open_websocket_connection", fake_open)
result = Runner.run_streamed(Agent(name="test", model=model), "hi")
async for _ in result.stream_events():
pass
assert result.context_wrapper.usage.requests == 1
assert result.context_wrapper.usage.total_tokens == 0
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_response_span_counts_request_without_usage() -> None:
class Responses:
async def create(self, **kwargs: Any) -> Response:
return _response_without_usage()
class Client:
responses = Responses()
base_url = httpx2.URL("https://custom.example.test/v1/")
model = OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, Client()))
with trace("test"):
await model.get_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.ENABLED,
)
spans = [span.export() for span in fetch_ordered_spans() if span.span_data.type == "response"]
assert len(spans) == 1
assert spans[0]["span_data"]["usage"]["requests"] == 1 # type: ignore[index]
assert spans[0]["span_data"]["usage"]["total_tokens"] == 0 # type: ignore[index]
@pytest.mark.allow_call_model_methods
@pytest.mark.asyncio
async def test_streamed_span_counts_request_before_terminal_event_close() -> None:
client = _streaming_client_for([_completed_event_without_usage()])
model = OpenAIResponsesModel(model="gpt-4", openai_client=cast(Any, client))
with trace("test"):
stream = model.stream_response(
system_instructions=None,
input="hi",
model_settings=ModelSettings(),
tools=[],
output_schema=None,
handoffs=[],
tracing=ModelTracing.ENABLED,
)
async for event in stream:
if event.type == "response.completed":
break
await stream.aclose()
spans = [span.export() for span in fetch_ordered_spans() if span.span_data.type == "response"]
assert len(spans) == 1
assert spans[0]["span_data"]["usage"]["requests"] == 1 # type: ignore[index]
assert spans[0]["span_data"]["usage"]["total_tokens"] == 0 # type: ignore[index]