mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(chatgpt): preserve streamed output in chat bridge
This commit is contained in:
parent
303434d573
commit
9091c66795
3 changed files with 115 additions and 21 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue