From 9091c66795b4aa1b915a0a29666bfe56c6aa2def Mon Sep 17 00:00:00 2001 From: Daniel Phang Date: Sat, 26 Sep 2026 20:28:31 -0700 Subject: [PATCH] fix(chatgpt): preserve streamed output in chat bridge --- .../handler.py | 52 +++++++++++---- .../transformation.py | 21 ++++--- ...itellm_responses_transformation_handler.py | 63 ++++++++++++++++++- 3 files changed, 115 insertions(+), 21 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 642a78789b2..ebbaec8c1bb 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -2,12 +2,13 @@ Handler for transforming /chat/completions api requests to litellm.responses requests """ -from collections.abc import AsyncIterable, Coroutine, Iterable +from collections.abc import AsyncIterable, Coroutine, Iterable, Mapping from typing import TYPE_CHECKING, Any, Final, Union +from pydantic import BaseModel from typing_extensions import TypedDict -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStreamEvents if TYPE_CHECKING: from litellm import CustomStreamWrapper, LiteLLMLoggingObj, ModelResponse @@ -54,17 +55,32 @@ class ResponsesToCompletionBridgeHandler: def _coerce_response_object( response_obj: object, hidden_params: dict | None, + stream_events: Iterable[object] = (), ) -> "ResponsesAPIResponse": if isinstance(response_obj, ResponsesAPIResponse): - response = response_obj + base_response = response_obj elif isinstance(response_obj, dict): try: - response = ResponsesAPIResponse(**response_obj) + base_response = ResponsesAPIResponse(**response_obj) except Exception: - response = ResponsesAPIResponse.model_construct(**response_obj) + base_response = ResponsesAPIResponse.model_construct(**response_obj) else: raise ValueError("Unexpected responses stream payload") + from .transformation import LiteLLMResponsesTransformationHandler + + parsed_chunks: Final = tuple( + payload + for event in stream_events + if (payload := ResponsesToCompletionBridgeHandler._stream_event_payload(event)) is not None + ) + recovered_output: Final = LiteLLMResponsesTransformationHandler.recover_output_items_from_chunks(parsed_chunks) + response: Final = ( + base_response + if base_response.output or not recovered_output + else base_response.model_copy(update={"output": recovered_output}) + ) + if hidden_params: existing: Final = getattr(response, "_hidden_params", None) if not isinstance(existing, dict) or not existing: @@ -74,9 +90,24 @@ class ResponsesToCompletionBridgeHandler: existing.setdefault(key, value) return response + @staticmethod + def _stream_event_payload(event: object) -> Mapping[str, object] | None: + if isinstance(event, Mapping): + return event + if isinstance(event, BaseModel): + return event.model_dump() + return None + + @staticmethod + def _is_recoverable_stream_event(event: object) -> bool: + event_type: Final = event.get("type") if isinstance(event, Mapping) else getattr(event, "type", None) + return event_type in ( + ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + ) + def _collect_response_from_stream(self, stream_iter: Iterable[object]) -> "ResponsesAPIResponse": - for _ in stream_iter: - pass + stream_events: Final = tuple(event for event in stream_iter if self._is_recoverable_stream_event(event)) completed: Final[object] = getattr(stream_iter, "completed_response", None) response_obj: Final[object] = getattr(completed, "response", None) if completed else None @@ -84,14 +115,13 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Stream ended without a completed response") hidden_params: Final = getattr(stream_iter, "_hidden_params", None) - response: Final = self._coerce_response_object(response_obj, hidden_params) + response: Final = self._coerce_response_object(response_obj, hidden_params, stream_events) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response async def _collect_response_from_stream_async(self, stream_iter: AsyncIterable[object]) -> "ResponsesAPIResponse": - async for _ in stream_iter: - pass + stream_events: Final = tuple([event async for event in stream_iter if self._is_recoverable_stream_event(event)]) completed: Final[object] = getattr(stream_iter, "completed_response", None) response_obj: Final[object] = getattr(completed, "response", None) if completed else None @@ -99,7 +129,7 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Stream ended without a completed response") hidden_params: Final = getattr(stream_iter, "_hidden_params", None) - response: Final = self._coerce_response_object(response_obj, hidden_params) + response: Final = self._coerce_response_object(response_obj, hidden_params, stream_events) if not isinstance(response, ResponsesAPIResponse): raise ValueError("Stream completed response is invalid") return response diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 31af5a144eb..844da405a98 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -833,18 +833,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return cast(list[dict[str, object]], response_output) @classmethod - def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]: - if not raw_sse or not isinstance(raw_sse, str): - return [] - + def recover_output_items_from_chunks(cls, parsed_chunks: Iterable[Mapping[str, object]]) -> list[dict[str, object]]: recovered_output_items: Final[dict[int, dict[str, object]]] = {} recovered_text_only_items: Final[dict[int, dict[str, object]]] = {} - for chunk in raw_sse.splitlines(): - parsed_chunk = parse_sse_json_chunk(chunk) - if parsed_chunk is None: - continue - + for parsed_chunk in parsed_chunks: event_type = parsed_chunk.get("type") if event_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED: @@ -881,6 +874,16 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return [] + @classmethod + def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]: + if not raw_sse or not isinstance(raw_sse, str): + return [] + + parsed_chunks: Final = tuple( + parsed_chunk for chunk in raw_sse.splitlines() if (parsed_chunk := parse_sse_json_chunk(chunk)) is not None + ) + return cls.recover_output_items_from_chunks(parsed_chunks) + @classmethod def _recover_output_items_from_logging(cls, logging_obj: "LiteLLMLoggingObj") -> list[dict[str, object]]: model_call_details: Final = getattr(logging_obj, "model_call_details", {}) or {} diff --git a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py index c5d7ca96a21..76dd5193187 100644 --- a/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py +++ b/tests/unit/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py @@ -1,15 +1,16 @@ from datetime import datetime +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch import pytest - import litellm from litellm.completion_extras.litellm_responses_transformation.handler import ( ResponsesToCompletionBridgeHandler, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse @@ -218,6 +219,66 @@ def _completed_chat_response() -> ModelResponse: ) +def _empty_responses_response() -> ResponsesAPIResponse: + return ResponsesAPIResponse.model_construct(output=[], error=None) + + +def _output_item_done_event() -> dict: + return { + "type": "response.output_item.done", + "output_index": 0, + "item": { + "type": "message", + "id": "msg_from_stream", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "Recovered from stream", + "annotations": [], + } + ], + }, + } + + +class _SyncResponsesStream: + def __init__(self, response: ResponsesAPIResponse): + self.completed_response = SimpleNamespace(response=response) + self._hidden_params = {} + + def __iter__(self): + return iter((_output_item_done_event(), self.completed_response)) + + +class _AsyncResponsesStream: + def __init__(self, response: ResponsesAPIResponse): + self.completed_response = SimpleNamespace(response=response) + self._hidden_params = {} + + async def __aiter__(self): + for event in (_output_item_done_event(), self.completed_response): + yield event + + +def test_collect_response_from_stream_recovers_output_items(): + bridge = ResponsesToCompletionBridgeHandler() + + response = bridge._collect_response_from_stream(_SyncResponsesStream(_empty_responses_response())) + + assert response.output == [_output_item_done_event()["item"]] + + +@pytest.mark.asyncio +async def test_collect_response_from_async_stream_recovers_output_items(): + bridge = ResponsesToCompletionBridgeHandler() + + response = await bridge._collect_response_from_stream_async(_AsyncResponsesStream(_empty_responses_response())) + + assert response.output == [_output_item_done_event()["item"]] + + @pytest.mark.asyncio async def test_acompletion_streams_completed_model_response(): """A streaming request whose bridge call comes back already completed must still be