mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(responses): tolerate dict response payloads in streaming events
This commit is contained in:
parent
24123269cc
commit
e343e3a1ab
9 changed files with 231 additions and 17 deletions
|
|
@ -1311,8 +1311,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
|
|||
|
||||
# Check if response contains function_call items in output
|
||||
# to determine correct finish_reason
|
||||
response_data = parsed_chunk.get("response", {})
|
||||
output_items = response_data.get("output", []) if response_data else []
|
||||
response_payload = parsed_chunk.get("response")
|
||||
response_data: dict[str, Any] = response_payload if isinstance(response_payload, dict) else {}
|
||||
output_items = response_data.get("output") or []
|
||||
|
||||
has_function_calls = any(
|
||||
item.get("type") == "function_call" for item in output_items if isinstance(item, dict)
|
||||
|
|
|
|||
|
|
@ -78,7 +78,10 @@ from litellm.litellm_core_utils.redact_messages import (
|
|||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
from litellm.responses.utils import (
|
||||
ResponseAPILoggingUtils,
|
||||
normalize_response_api_usage,
|
||||
)
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
from litellm.types.containers.main import ContainerObject
|
||||
from litellm.types.llms.openai import (
|
||||
|
|
@ -3224,13 +3227,17 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
(ResponseCompletedEvent, ResponseIncompleteEvent, ResponseFailedEvent),
|
||||
):
|
||||
## return unified Usage object
|
||||
if isinstance(result.response.usage, ResponseAPIUsage):
|
||||
response = result.response
|
||||
if isinstance(response, dict):
|
||||
response = ResponsesAPIResponse.model_construct(**response)
|
||||
responses_api_usage = normalize_response_api_usage(response.usage)
|
||||
if responses_api_usage is not None:
|
||||
transformed_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
|
||||
result.response.usage
|
||||
responses_api_usage
|
||||
)
|
||||
# Set as dict instead of Usage object so model_dump() serializes it correctly
|
||||
setattr(
|
||||
result.response,
|
||||
response,
|
||||
"usage",
|
||||
(
|
||||
transformed_usage.model_dump()
|
||||
|
|
@ -3238,7 +3245,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
else dict(transformed_usage)
|
||||
),
|
||||
)
|
||||
return result.response
|
||||
return response
|
||||
else:
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -30,6 +30,23 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
def _with_typed_response_payload(parsed_chunk: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Build the nested ``ResponsesAPIResponse`` so ``model_construct`` doesn't leave it a raw dict."""
|
||||
from litellm.responses.utils import normalize_response_api_usage
|
||||
|
||||
response_payload = parsed_chunk.get("response")
|
||||
if not isinstance(response_payload, dict):
|
||||
return parsed_chunk
|
||||
try:
|
||||
response = ResponsesAPIResponse(**response_payload)
|
||||
except ValidationError:
|
||||
response = ResponsesAPIResponse.model_construct(**response_payload)
|
||||
normalized_usage = normalize_response_api_usage(response.usage)
|
||||
if normalized_usage is not None:
|
||||
response.usage = normalized_usage
|
||||
return {**parsed_chunk, "response": response}
|
||||
|
||||
|
||||
class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
|
|
@ -349,7 +366,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
|
|||
event_pydantic_model.__name__,
|
||||
parsed_chunk,
|
||||
)
|
||||
return event_pydantic_model.model_construct(**parsed_chunk)
|
||||
return event_pydantic_model.model_construct(**_with_typed_response_payload(parsed_chunk))
|
||||
|
||||
@staticmethod
|
||||
def get_event_model_class(event_type: str) -> Any:
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from typing import (
|
|||
Iterable,
|
||||
List,
|
||||
Mapping,
|
||||
MutableMapping,
|
||||
Optional,
|
||||
Type,
|
||||
Union,
|
||||
|
|
@ -14,7 +15,7 @@ from typing import (
|
|||
overload,
|
||||
)
|
||||
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -1010,6 +1011,42 @@ class ResponsesAPIRequestUtils:
|
|||
)
|
||||
|
||||
|
||||
def normalize_response_api_usage(usage: object) -> ResponseAPIUsage | None:
|
||||
"""Coerce a usage payload to ``ResponseAPIUsage``; returns None when it isn't Responses-API usage.
|
||||
|
||||
Streaming events built via ``model_construct`` (provider payloads that fail validation) carry
|
||||
their usage as a plain dict, so attribute access on it would raise.
|
||||
"""
|
||||
if isinstance(usage, ResponseAPIUsage):
|
||||
return usage
|
||||
if isinstance(usage, Mapping):
|
||||
try:
|
||||
return ResponseAPIUsage(**dict(usage))
|
||||
except ValidationError:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def get_response_api_usage(
|
||||
response: ResponsesAPIResponse | Mapping[str, Any],
|
||||
) -> ResponseAPIUsage | None:
|
||||
"""Read usage off a Responses API response payload that may still be an unvalidated dict."""
|
||||
if isinstance(response, Mapping):
|
||||
return normalize_response_api_usage(response.get("usage"))
|
||||
return normalize_response_api_usage(response.usage)
|
||||
|
||||
|
||||
def set_response_api_usage(
|
||||
response: ResponsesAPIResponse | MutableMapping[str, Any],
|
||||
usage: ResponseAPIUsage,
|
||||
) -> None:
|
||||
"""Write usage back onto a Responses API response payload that may still be an unvalidated dict."""
|
||||
if isinstance(response, MutableMapping):
|
||||
response["usage"] = usage
|
||||
else:
|
||||
response.usage = usage
|
||||
|
||||
|
||||
class ResponseAPILoggingUtils:
|
||||
@staticmethod
|
||||
def _is_response_api_usage(usage: Union[dict, ResponseAPIUsage]) -> bool:
|
||||
|
|
|
|||
|
|
@ -2242,7 +2242,9 @@ class Router:
|
|||
completed,
|
||||
(ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent),
|
||||
):
|
||||
return completed.response.usage
|
||||
from litellm.responses.utils import get_response_api_usage
|
||||
|
||||
return get_response_api_usage(completed.response)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -2264,6 +2266,10 @@ class Router:
|
|||
and produce a clean ResponseAPIUsage — no token-naming split, no
|
||||
setattr bypass.
|
||||
"""
|
||||
from litellm.responses.utils import (
|
||||
get_response_api_usage,
|
||||
set_response_api_usage,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseAPIUsage,
|
||||
ResponseCompletedEvent,
|
||||
|
|
@ -2276,15 +2282,17 @@ class Router:
|
|||
(ResponseCompletedEvent, ResponseFailedEvent, ResponseIncompleteEvent),
|
||||
):
|
||||
return
|
||||
response = fallback_item.response
|
||||
if response.usage is None:
|
||||
fb = get_response_api_usage(fallback_item.response)
|
||||
if fb is None:
|
||||
return
|
||||
|
||||
fb = response.usage
|
||||
response.usage = ResponseAPIUsage(
|
||||
input_tokens=(partial_usage.input_tokens or 0) + (fb.input_tokens or 0),
|
||||
output_tokens=(partial_usage.output_tokens or 0) + (fb.output_tokens or 0),
|
||||
total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0),
|
||||
set_response_api_usage(
|
||||
fallback_item.response,
|
||||
ResponseAPIUsage(
|
||||
input_tokens=(partial_usage.input_tokens or 0) + (fb.input_tokens or 0),
|
||||
output_tokens=(partial_usage.output_tokens or 0) + (fb.output_tokens or 0),
|
||||
total_tokens=(partial_usage.total_tokens or 0) + (fb.total_tokens or 0),
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -107,6 +107,48 @@ def test_combine_responses_fallback_usage_sums_completed_event():
|
|||
assert combined.total_tokens == 26
|
||||
|
||||
|
||||
def _make_completed_event_with_dict_response(
|
||||
input_tokens: int, output_tokens: int, total_tokens: int
|
||||
) -> ResponseCompletedEvent:
|
||||
"""A provider payload that failed validation keeps `response` as a raw dict (model_construct)."""
|
||||
return ResponseCompletedEvent.model_construct(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response={
|
||||
"id": "resp_123",
|
||||
"usage": {
|
||||
"input_tokens": input_tokens,
|
||||
"output_tokens": output_tokens,
|
||||
"total_tokens": total_tokens,
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def test_extract_partial_responses_usage_dict_response():
|
||||
"""Regression test for https://github.com/BerriAI/litellm/issues/34754: dict response payload."""
|
||||
source = MagicMock()
|
||||
source.completed_response = _make_completed_event_with_dict_response(11, 7, 18)
|
||||
|
||||
usage = Router._extract_partial_responses_usage(source)
|
||||
assert usage is not None
|
||||
assert (usage.input_tokens, usage.output_tokens, usage.total_tokens) == (11, 7, 18)
|
||||
|
||||
|
||||
def test_combine_responses_fallback_usage_dict_response():
|
||||
"""Usage is summed into a fallback event whose response payload is a raw dict."""
|
||||
fallback_event = _make_completed_event_with_dict_response(5, 3, 8)
|
||||
partial = ResponseAPIUsage(input_tokens=11, output_tokens=7, total_tokens=18)
|
||||
|
||||
Router._combine_responses_fallback_usage(fallback_event, partial)
|
||||
|
||||
combined = fallback_event.response["usage"]
|
||||
assert (combined.input_tokens, combined.output_tokens, combined.total_tokens) == (
|
||||
16,
|
||||
10,
|
||||
26,
|
||||
)
|
||||
|
||||
|
||||
def test_combine_responses_fallback_usage_passthrough_for_unknown_event():
|
||||
"""Events that are not completed/failed/incomplete are not mutated."""
|
||||
other = MagicMock() # not a ResponseCompletedEvent etc. → isinstance false
|
||||
|
|
|
|||
|
|
@ -1075,6 +1075,40 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason()
|
|||
), "response.completed with function_call output should emit finish_reason='tool_calls'"
|
||||
|
||||
|
||||
def test_response_completed_with_null_output_emits_stop_finish_reason():
|
||||
"""
|
||||
Regression test for https://github.com/BerriAI/litellm/issues/34754
|
||||
|
||||
Upstream providers may send `response.output: null` on response.completed; iterating that
|
||||
None used to raise `TypeError`/`AttributeError` while bridging /responses back to
|
||||
/chat/completions.
|
||||
"""
|
||||
from litellm.completion_extras.litellm_responses_transformation.transformation import (
|
||||
OpenAiResponsesToChatCompletionStreamIterator,
|
||||
)
|
||||
|
||||
iterator = OpenAiResponsesToChatCompletionStreamIterator(
|
||||
streaming_response=None, sync_stream=True
|
||||
)
|
||||
|
||||
chunk = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"status": "completed",
|
||||
"output": None,
|
||||
"usage": {"input_tokens": 5, "output_tokens": 3, "total_tokens": 8},
|
||||
},
|
||||
}
|
||||
|
||||
result = iterator.chunk_parser(chunk)
|
||||
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
assert result.usage is not None
|
||||
assert result.usage.prompt_tokens == 5
|
||||
assert result.usage.completion_tokens == 3
|
||||
|
||||
|
||||
def test_response_completed_with_message_only_emits_stop_finish_reason():
|
||||
"""
|
||||
Test that response.completed with only message output (no function_call) emits finish_reason='stop'.
|
||||
|
|
|
|||
|
|
@ -2678,6 +2678,44 @@ def test_streaming_success_handler_includes_vertex_ai_metadata_in_standard_loggi
|
|||
assert payload["response"]["vertex_ai_url_context_metadata"] == url_context_metadata
|
||||
|
||||
|
||||
def test_get_assembled_streaming_response_handles_dict_response_payload():
|
||||
"""Regression test for https://github.com/BerriAI/litellm/issues/34754
|
||||
|
||||
Completed events built via model_construct carry a raw dict response; usage must still be
|
||||
assembled into chat-style usage instead of raising AttributeError and dropping the log.
|
||||
"""
|
||||
import datetime
|
||||
|
||||
from litellm.types.llms.openai import (
|
||||
ResponseCompletedEvent,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
|
||||
logging_obj = _make_logging_obj(stream=True)
|
||||
event = ResponseCompletedEvent.model_construct(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response={
|
||||
"id": "resp-1",
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"usage": {"input_tokens": 9, "output_tokens": 6, "total_tokens": 15},
|
||||
},
|
||||
)
|
||||
|
||||
assembled = logging_obj._get_assembled_streaming_response(
|
||||
result=event,
|
||||
start_time=datetime.datetime.now(),
|
||||
end_time=datetime.datetime.now(),
|
||||
is_async=True,
|
||||
streaming_chunks=[],
|
||||
)
|
||||
|
||||
assert assembled is not None
|
||||
assert assembled.usage["prompt_tokens"] == 9
|
||||
assert assembled.usage["completion_tokens"] == 6
|
||||
assert assembled.usage["total_tokens"] == 15
|
||||
|
||||
|
||||
def test_get_assembled_streaming_response_returns_none_for_non_streaming_text_completion():
|
||||
"""Non-streaming TextCompletionResponse should also return None."""
|
||||
import datetime
|
||||
|
|
|
|||
|
|
@ -288,6 +288,36 @@ class TestOpenAIResponsesAPIConfig:
|
|||
assert result.type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
assert result.response.id == "resp_123"
|
||||
|
||||
def test_transform_streaming_response_invalid_payload_keeps_typed_response(self):
|
||||
"""Regression test for https://github.com/BerriAI/litellm/issues/34754
|
||||
|
||||
A response payload that fails validation falls back to model_construct; the nested
|
||||
response must still be a ResponsesAPIResponse so `event.response.usage` works.
|
||||
"""
|
||||
completed_chunk = {
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_123",
|
||||
"status": "completed",
|
||||
"output": [],
|
||||
"usage": {
|
||||
"input_tokens": 12,
|
||||
"output_tokens": 4,
|
||||
"total_tokens": 16,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
result = self.config.transform_streaming_response(
|
||||
model=self.model,
|
||||
parsed_chunk=completed_chunk,
|
||||
logging_obj=self.logging_obj,
|
||||
)
|
||||
|
||||
assert isinstance(result, ResponseCompletedEvent)
|
||||
assert isinstance(result.response, ResponsesAPIResponse)
|
||||
assert result.response.usage.input_tokens == 12
|
||||
|
||||
@pytest.mark.serial
|
||||
def test_validate_environment(self):
|
||||
"""Test that validate_environment correctly sets the Authorization header"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue