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
85 changes: 77 additions & 8 deletions src/agents/extensions/models/any_llm_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import inspect
import json
import time
from collections.abc import AsyncGenerator, AsyncIterator, Iterable
from collections.abc import AsyncGenerator, AsyncIterator, Iterable, Mapping
from copy import copy
from typing import TYPE_CHECKING, Any, Literal, cast, overload

Expand Down Expand Up @@ -52,7 +52,12 @@
from ...tracing import generation_span, response_span
from ...tracing.span_data import GenerationSpanData
from ...tracing.spans import Span
from ...usage import Usage
from ...usage import (
Usage,
_attach_raw_usage_snapshot,
_extract_raw_usage_snapshot,
_raw_usage_snapshot,
)
from ...util._error_tracing import model_span_errors, record_model_error_on_span
from ...util._json import _to_dump_compatible

Expand All @@ -75,6 +80,12 @@ class InternalChatCompletionMessage(ChatCompletionMessage):
reasoning_content: str = ""


def _usage_payload(response: Any) -> Any | None:
if isinstance(response, Mapping):
return response.get("usage")
return getattr(response, "usage", None)


class _AnyLLMResponsesParamsShim:
"""Fallback shim for tests and older any-llm layouts."""

Expand Down Expand Up @@ -417,6 +428,11 @@ async def _get_response_via_responses(
usage=usage,
response_id=response.id,
request_id=getattr(response, "_request_id", None),
raw_usage=(
_extract_raw_usage_snapshot(response, fallback=response.usage)
if model_settings.preserve_raw_usage is True
else None
),
)

async def _stream_response_via_responses(
Expand Down Expand Up @@ -463,6 +479,8 @@ async def _stream_response_via_responses(
chunk_type = getattr(chunk, "type", None)
if isinstance(chunk, ResponseCompletedEvent):
final_response = chunk.response
if model_settings.preserve_raw_usage is True:
_attach_raw_usage_snapshot(chunk.response, chunk.response.usage)
elif chunk_type in {"response.failed", "response.incomplete"}:
terminal_response = getattr(chunk, "response", None)
terminal_failure_error = response_terminal_failure_error(
Expand Down Expand Up @@ -640,7 +658,16 @@ async def _get_response_via_chat(
if logprob_models:
self._attach_logprobs_to_output(items, logprob_models)

return ModelResponse(output=items, usage=usage, response_id=None)
return ModelResponse(
output=items,
usage=usage,
response_id=None,
raw_usage=(
_extract_raw_usage_snapshot(response, fallback=response.usage)
if model_settings.preserve_raw_usage is True
else None
),
)

async def _stream_response_via_chat(
self,
Expand Down Expand Up @@ -686,11 +713,21 @@ async def _stream_response_via_chat(
final_response: Response | None = None
yielded_terminal_event = False
close_stream_in_background = False
raw_usage_options: dict[str, Any] = (
{"preserve_raw_usage": True} if model_settings.preserve_raw_usage is True else {}
)
try:
async for chunk in ChatCmplStreamHandler.handle_stream(
response,
cast(Any, self._normalize_chat_stream(stream)),
cast(
Any,
self._normalize_chat_stream(
stream,
preserve_raw_usage=model_settings.preserve_raw_usage is True,
),
),
model=self.model,
**raw_usage_options,
):
# Record terminal state and populate the span before yielding so a consumer
# that stops at the completed event still leaves a fully recorded span.
Expand Down Expand Up @@ -899,7 +936,18 @@ async def _fetch_chat_response(
)

if not stream:
return self._normalize_chat_completion_response(ret)
raw_usage = (
_raw_usage_snapshot(_usage_payload(ret))
if model_settings.preserve_raw_usage is True
else None
)
normalized_response = self._normalize_chat_completion_response(ret)
if model_settings.preserve_raw_usage is True:
_attach_raw_usage_snapshot(
normalized_response,
raw_usage,
)
return normalized_response

responses_tool_choice = OpenAIResponsesConverter.convert_tool_choice(
model_settings.tool_choice
Expand Down Expand Up @@ -1051,7 +1099,18 @@ async def _fetch_responses_response(
if stream:
return cast(AsyncIterator[ResponseStreamEvent], response)

return self._normalize_response(response)
raw_usage = (
_raw_usage_snapshot(_usage_payload(response))
if model_settings.preserve_raw_usage is True
else None
)
normalized_response = self._normalize_response(response)
if model_settings.preserve_raw_usage is True:
_attach_raw_usage_snapshot(
normalized_response,
raw_usage,
)
return normalized_response

@staticmethod
def _split_model_name(model: str) -> tuple[str, str]:
Expand Down Expand Up @@ -1183,10 +1242,20 @@ def _normalize_chat_completion_response(self, response: Any) -> ChatCompletion:
return ChatCompletion.model_validate(response)

async def _normalize_chat_stream(
self, stream: AsyncIterator[ChatCompletionChunk]
self,
stream: AsyncIterator[ChatCompletionChunk],
*,
preserve_raw_usage: bool = False,
) -> AsyncIterator[ChatCompletionChunk]:
async for chunk in stream:
yield self._normalize_chat_chunk(chunk)
raw_usage = _raw_usage_snapshot(_usage_payload(chunk)) if preserve_raw_usage else None
normalized_chunk = self._normalize_chat_chunk(chunk)
if preserve_raw_usage:
_attach_raw_usage_snapshot(
normalized_chunk,
raw_usage,
)
yield normalized_chunk

def _normalize_chat_chunk(self, chunk: Any) -> ChatCompletionChunk:
normalized_chunk = chunk
Expand Down
9 changes: 9 additions & 0 deletions src/agents/items.py
Original file line number Diff line number Diff line change
Expand Up @@ -678,6 +678,15 @@ class ModelResponse:
request_id: str | None = None
"""The transport request ID for this model call, if provided by the model SDK."""

raw_usage: dict[str, Any] | None = None
"""A JSON-compatible snapshot of the provider usage payload, when preservation is enabled.

The snapshot is captured only while the unnormalized provider payload is available, before the
Agents SDK normalizes missing usage fields. It is ``None`` when preservation is disabled, no
usage payload reaches the model adapter, or upstream normalization has already discarded
field-presence information.
"""

def to_input_items(self) -> list[TResponseInputItem]:
"""Convert the output into a list of input items suitable for passing to the model."""
# Most output items can be replayed via a direct model_dump. Tool-search items carry
Expand Down
11 changes: 11 additions & 0 deletions src/agents/model_settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,16 @@ class ModelSettings:
control which prompt prefixes are eligible for caching.
"""

preserve_raw_usage: bool | None = None
"""Whether to preserve the provider usage payload on completed model responses.

When enabled and the model adapter still has the unnormalized provider payload,
``ModelResponse.raw_usage`` contains a JSON-compatible snapshot captured before the Agents
SDK normalizes missing usage fields. It remains ``None`` when usage is absent or upstream
normalization has already discarded field-presence information. This setting does not request
usage from the provider; use ``include_usage`` separately when a streaming provider requires it.
"""

if TYPE_CHECKING:

def __init__(
Expand Down Expand Up @@ -228,6 +238,7 @@ def __init__(
retry: ModelRetrySettings | dict[str, Any] | None = None,
context_management: list[ContextManagement] | None = None,
prompt_cache_options: PromptCacheOptions | None = None,
preserve_raw_usage: bool | None = None,
) -> None: ...

def resolve(self, override: ModelSettings | dict[str, Any] | None) -> ModelSettings:
Expand Down
15 changes: 14 additions & 1 deletion src/agents/models/chatcmpl_stream_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -51,7 +51,12 @@
from ..exceptions import ModelBehaviorError, UserError
from ..items import TResponseStreamEvent
from ..logger import logger
from ..usage import _cache_write_tokens, _make_input_tokens_details
from ..usage import (
_attach_raw_usage_snapshot,
_cache_write_tokens,
_extract_raw_usage_snapshot,
_make_input_tokens_details,
)
from .chatcmpl_helpers import ChatCmplHelpers
from .fake_id import FAKE_RESPONSES_ID

Expand Down Expand Up @@ -584,6 +589,7 @@ async def handle_stream(
stream: AsyncStream[ChatCompletionChunk],
model: str | None = None,
strict_feature_validation: bool = False,
preserve_raw_usage: bool = False,
) -> AsyncIterator[TResponseStreamEvent]:
"""
Handle a streaming chat completion response and yield response events.
Expand All @@ -593,8 +599,11 @@ async def handle_stream(
stream: The async stream of chat completion chunks from the model
model: The source model that is generating this stream. Used to handle
provider-specific stream processing.
preserve_raw_usage: Whether to retain the last provider usage payload before
converting it to the Responses usage shape.
"""
usage: CompletionUsage | None = None
raw_usage: dict[str, Any] | None = None
state = StreamingState()
output_layout = _StreamOutputLayout()
sequence_number = SequenceNumber()
Expand All @@ -616,6 +625,8 @@ async def handle_stream(
# Only update when chunk has usage data (not always in the last chunk)
if hasattr(chunk, "usage") and chunk.usage is not None:
usage = chunk.usage
if preserve_raw_usage:
raw_usage = _extract_raw_usage_snapshot(chunk, fallback=chunk.usage)

if not chunk.choices:
continue
Expand Down Expand Up @@ -1272,6 +1283,8 @@ async def handle_stream(
if usage
else None
)
if preserve_raw_usage:
_attach_raw_usage_snapshot(final_response, raw_usage)

yield ResponseCompletedEvent(
response=final_response,
Expand Down
11 changes: 10 additions & 1 deletion src/agents/models/openai_chatcompletions.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,7 @@
from ..tracing import generation_span
from ..tracing.span_data import GenerationSpanData
from ..tracing.spans import Span
from ..usage import Usage
from ..usage import Usage, _raw_usage_snapshot
from ..util._error_tracing import model_span_errors
from ..util._json import _to_dump_compatible
from ._openai_retry import get_openai_retry_advice
Expand Down Expand Up @@ -351,6 +351,11 @@ async def get_response(
# The OpenAI SDK records the `x-request-id` header on every parsed response,
# so callers can inspect the same debugging handle as on the Responses path.
request_id=getattr(response, "_request_id", None),
raw_usage=(
_raw_usage_snapshot(response.usage)
if model_settings.preserve_raw_usage is True
else None
),
)

@staticmethod
Expand Down Expand Up @@ -444,6 +449,9 @@ async def stream_response(
else:
stream_for_handler = stream

raw_usage_options: dict[str, Any] = (
{"preserve_raw_usage": True} if model_settings.preserve_raw_usage is True else {}
)
close_stream_in_background = False
yielded_terminal_event = False
try:
Expand All @@ -452,6 +460,7 @@ async def stream_response(
cast(AsyncStream[ChatCompletionChunk], stream_for_handler),
model=self.model,
strict_feature_validation=self._strict_feature_validation,
**raw_usage_options,
):
if chunk.type == "response.completed":
final_response = chunk.response
Expand Down
15 changes: 14 additions & 1 deletion src/agents/models/openai_responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,13 @@
validate_responses_tool_search_configuration,
)
from ..tracing import SpanError, response_span
from ..usage import Usage, _response_usage_to_usage, model_usage_to_span_usage
from ..usage import (
Usage,
_attach_raw_usage_snapshot,
_raw_usage_snapshot,
_response_usage_to_usage,
model_usage_to_span_usage,
)
from ..util._error_tracing import record_model_error_on_span
from ..util._json import _to_dump_compatible
from ..version import __version__
Expand Down Expand Up @@ -550,6 +556,11 @@ async def get_response(
usage=usage,
response_id=response.id,
request_id=getattr(response, "_request_id", None),
raw_usage=(
_raw_usage_snapshot(response.usage)
if model_settings.preserve_raw_usage is True
else None
),
)

async def stream_response(
Expand Down Expand Up @@ -592,6 +603,8 @@ async def stream_response(
chunk_type = getattr(chunk, "type", None)
if isinstance(chunk, ResponseCompletedEvent):
final_response = chunk.response
if model_settings.preserve_raw_usage is True:
_attach_raw_usage_snapshot(chunk.response, chunk.response.usage)
elif chunk_type in {
"response.failed",
"response.incomplete",
Expand Down
7 changes: 6 additions & 1 deletion src/agents/run_internal/run_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,7 @@
from ..tracing.config import include_task_and_turn_spans
from ..tracing.model_tracing import get_model_tracing_impl
from ..tracing.span_data import AgentSpanData, TaskSpanData
from ..usage import Usage, _response_usage_to_usage
from ..usage import Usage, _extract_raw_usage_snapshot, _response_usage_to_usage
from ..util import _coro, _error_tracing
from ..util._asyncio_tasks import gather_with_cancel
from .agent_bindings import AgentBindings, bind_public_agent
Expand Down Expand Up @@ -1774,6 +1774,11 @@ async def rewind_model_request() -> None:
usage=usage,
response_id=terminal_response.id,
request_id=getattr(terminal_response, "_request_id", None),
raw_usage=(
_extract_raw_usage_snapshot(terminal_response)
if model_settings.preserve_raw_usage is True
else None
),
)

if isinstance(event, ResponseOutputItemDoneEvent):
Expand Down
Loading