fix(bridge): carry provider metadata on streamed chats and keep served ids in spend logs

The terminal chunk of a bridged streaming chat now carries the same provider
fields the non-streaming response does (service_tier, content_filters), so
streaming clients see them on the final chunk

The passthrough drops Responses API bookkeeping (background, top_logprobs,
store, ...) by subtracting the OpenAI SDK's Response schema from the fields it
copies instead of a hand-kept denylist

Spend log rows for /v1/messages calls served through the Responses adapter keep
the id the client was handed instead of the decoded upstream id
This commit is contained in:
mateo-berri 2026-09-05 19:21:40 -07:00
parent 089b4b8ff4
commit 64cbe6d0aa
4 changed files with 162 additions and 10 deletions

View file

@ -5,8 +5,11 @@ Handler for transforming /chat/completions api requests to litellm.responses req
import json
import os
from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args
from openai.types.chat import ChatCompletion
from openai.types.responses import Response
from openai.types.responses.custom_tool_param import CustomToolParam
from openai.types.responses.response_input_param import (
FunctionCallOutput,
@ -43,6 +46,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolParamFunctionChunk,
Reasoning,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponsesAPIStreamEvents,
)
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
@ -70,6 +74,19 @@ if TYPE_CHECKING:
_CHAT_COMPLETION_FIELDS: Final = frozenset((*ModelResponse.model_fields, "usage"))
_RESPONSES_API_ONLY_FIELDS: Final = frozenset((*Response.model_fields, *ResponsesAPIResponse.model_fields)) - frozenset(
ChatCompletion.model_fields
)
def _provider_metadata(response_fields: Mapping[str, object] | None) -> Mapping[str, object]:
return MappingProxyType(
{
key: value
for key, value in (response_fields.items() if response_fields else ())
if value is not None and key not in _CHAT_COMPLETION_FIELDS and key not in _RESPONSES_API_ONLY_FIELDS
}
)
def _upstream_response_id(response_id: str | None) -> str | None:
@ -914,10 +931,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
)
model_response.id = _upstream_response_id(raw_response.id) or raw_response.id
provider_extras: Final = raw_response.model_extra.items() if raw_response.model_extra else ()
for key, value in provider_extras:
if key not in _CHAT_COMPLETION_FIELDS and value is not None:
setattr(model_response, key, value)
for key, value in _provider_metadata(raw_response.model_extra).items():
setattr(model_response, key, value)
# Preserve hidden params from the ResponsesAPIResponse, especially the headers
# which contain important provider information like x-request-id
@ -1551,6 +1566,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
from litellm.responses.utils import ResponseAPILoggingUtils
usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(response_data.get("usage"))
provider_metadata: Final = _provider_metadata(response_data)
return ModelResponseStream(
choices=[
StreamingChoices(
@ -1563,6 +1579,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
)
],
usage=usage,
provider_specific_fields=dict(provider_metadata) or None, # mutable-ok: field is typed dict
)
else:
pass

View file

@ -3918,11 +3918,12 @@ class Logging(LiteLLMLoggingBaseClass):
LiteLLMResponsesTransformationHandler,
)
served_id: Final = _provider_response_id(result)
try:
return LiteLLMResponsesTransformationHandler().transform_response(
translated: Final = LiteLLMResponsesTransformationHandler().transform_response(
model=self.model,
raw_response=result,
model_response=litellm.ModelResponse(id=_provider_response_id(result)),
model_response=litellm.ModelResponse(id=served_id),
logging_obj=self,
request_data={},
messages=[],
@ -3930,6 +3931,8 @@ class Logging(LiteLLMLoggingBaseClass):
litellm_params={},
encoding=litellm.encoding,
)
translated.id = served_id or translated.id
return translated
except Exception as e:
verbose_logger.debug(
"Responses API -> ModelResponse translation failed for "
@ -3937,7 +3940,7 @@ class Logging(LiteLLMLoggingBaseClass):
"usage-only ModelResponse to keep the spend_logs row.",
str(e),
)
model_response: Final = litellm.ModelResponse(id=_provider_response_id(result))
model_response: Final = litellm.ModelResponse(id=served_id)
model_response.model = self.model
usage: Final = getattr(result, "usage", None)
if usage is not None and ResponseAPILoggingUtils._is_response_api_usage(usage):

View file

@ -3955,10 +3955,11 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_tool
def _litellm_encoded_response_id(upstream_id: str) -> str:
import base64
from litellm.responses.utils import ResponsesAPIRequestUtils
tagged = f"litellm:custom_llm_provider:azure;model_id:deployment-1;response_id:{upstream_id}"
return "resp_" + base64.b64encode(tagged.encode()).decode()
return ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="azure", model_id="deployment-1", response_id=upstream_id
)
def test_transform_response_keeps_upstream_id_and_provider_extras():
@ -3997,6 +3998,9 @@ def test_transform_response_keeps_upstream_id_and_provider_extras():
"service_tier": "default",
"content_filters": content_filters,
"max_tool_calls": None,
"background": False,
"top_logprobs": 0,
"store": True,
}
)
model_response = ModelResponse(
@ -4029,9 +4033,63 @@ def test_transform_response_keeps_upstream_id_and_provider_extras():
assert "output" not in dumped and "status" not in dumped, (
"Responses schema fields must not leak into the chat response"
)
assert not {"background", "top_logprobs", "store"} & dumped.keys(), (
"Responses API bookkeeping must not ride along as chat metadata"
)
assert dumped["choices"][0]["message"]["tool_calls"][0]["function"]["name"] == "lookup_weather"
def test_bridged_response_is_priced_by_the_reported_service_tier():
from unittest.mock import Mock
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import ModelResponse
raw_response = ResponsesAPIResponse.model_validate(
{
"id": "resp_flex",
"created_at": 1734366691,
"object": "response",
"model": "gpt-5.4",
"status": "completed",
"output": [
{
"type": "message",
"id": "msg_1",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"usage": {"input_tokens": 1000, "output_tokens": 100, "total_tokens": 1100},
"service_tier": "flex",
}
)
result = LiteLLMResponsesTransformationHandler().transform_response(
model="gpt-5.4",
raw_response=raw_response,
model_response=ModelResponse(),
logging_obj=Mock(),
request_data={"model": "gpt-5.4"},
messages=[{"role": "user", "content": "hi"}],
optional_params={},
litellm_params={},
encoding=Mock(),
)
pricing = litellm.model_cost["gpt-5.4"]
flex_cost = 1000 * pricing["input_cost_per_token_flex"] + 100 * pricing["output_cost_per_token_flex"]
standard_cost = 1000 * pricing["input_cost_per_token"] + 100 * pricing["output_cost_per_token"]
cost = litellm.completion_cost(completion_response=result, custom_llm_provider="openai")
assert cost == pytest.approx(flex_cost)
assert cost < standard_cost
def test_streaming_chunks_carry_the_upstream_response_id():
from litellm.completion_extras.litellm_responses_transformation.transformation import (
OpenAiResponsesToChatCompletionStreamIterator,
@ -4048,3 +4106,44 @@ def test_streaming_chunks_carry_the_upstream_response_id():
ids = [iterator.chunk_parser(event).id for event in events]
assert ids == ["resp_azure_stream"] * len(events), f"streamed chunks did not carry the upstream id: {ids}"
def test_streaming_final_chunk_carries_provider_metadata():
from unittest.mock import MagicMock
from litellm.completion_extras.litellm_responses_transformation.transformation import (
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True)
content_filters = [{"blocked": False, "source_type": "completion", "content_filter_results": {}}]
events = [
{"type": "response.created", "response": {"id": "resp_azure_stream", "output": []}},
{"type": "response.output_text.delta", "delta": "Hello"},
{
"type": "response.completed",
"response": {
"id": "resp_azure_stream",
"output": [{"type": "message"}],
"usage": {"input_tokens": 3, "output_tokens": 1, "total_tokens": 4},
"service_tier": "default",
"content_filters": content_filters,
"background": False,
},
},
]
stream = CustomStreamWrapper(
completion_stream=iter([iterator.chunk_parser(event) for event in events]),
model="gpt-5.6",
custom_llm_provider="azure",
logging_obj=MagicMock(),
)
chunks = [chunk.model_dump() for chunk in stream]
assert chunks[-1]["choices"][0]["finish_reason"] == "stop"
assert chunks[-1]["service_tier"] == "default"
assert chunks[-1]["content_filters"] == content_filters
assert "background" not in chunks[-1]
assert all("service_tier" not in chunk for chunk in chunks[:-1])

View file

@ -4391,6 +4391,39 @@ def test_handle_anthropic_messages_response_logging_translates_bare_responses_ap
assert result.usage.total_tokens == 18 # type: ignore[attr-defined]
def test_handle_anthropic_messages_response_logging_keeps_the_served_response_id():
from openai.types.responses import ResponseOutputMessage, ResponseOutputText
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
served_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai", model_id="deployment-1", response_id="resp_upstream"
)
logging_obj = _anthropic_messages_logging_obj()
result = logging_obj._handle_anthropic_messages_response_logging(
result=ResponsesAPIResponse(
id=served_id,
created_at=1700000000,
output=[
ResponseOutputMessage(
id="msg-1",
type="message",
role="assistant",
status="completed",
content=[ResponseOutputText(annotations=[], text="hi", type="output_text")],
)
],
usage=ResponseAPIUsage(input_tokens=2, output_tokens=1, total_tokens=3),
service_tier="flex",
)
)
assert isinstance(result, ModelResponse)
assert result.id == served_id, "the spend log row must keep the id the caller was served"
assert result.service_tier == "flex"
def test_handle_anthropic_messages_response_logging_passes_model_response_through():
"""Anthropic-native path already yields a ModelResponse; it must be returned unchanged."""
logging_obj = _anthropic_messages_logging_obj()