diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 9f9016c5a7f..cac503b3ee5 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -6,13 +6,15 @@ import time import traceback import uuid from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass, replace from datetime import datetime from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, overload, runtime_checkable import httpx from openai._streaming import SSEDecoder +from pydantic import BaseModel from typing_extensions import TypeIs import litellm @@ -34,7 +36,19 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils from litellm.types.llms.openai import ( PART_UNION_TYPES, + BaseLiteLLMOpenAIResponseObject, + ContentPartAddedEvent, + ContentPartDoneEvent, + ContentPartDonePartOutputText, + ContentPartDonePartRefusal, + FunctionCallArgumentsDoneEvent, + OutputItemAddedEvent, + OutputItemDoneEvent, + OutputTextDoneEvent, + RefusalDoneEvent, ResponseAPIUsage, + ResponseCreatedEvent, + ResponseInProgressEvent, ResponsesAPIResponse, ResponsesAPIStreamEvents, ResponsesAPIStreamingResponse, @@ -220,6 +234,446 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None ) +def _obj_get(obj: object, key: str, default: object | None = None) -> object: + if obj is None: + return default + if isinstance(obj, Mapping): + source: Mapping[object, object] = obj + return source.get(key, default) + return getattr(obj, key, default) + + +def _safe_int(value: object, default: int) -> int: + if isinstance(value, bool): + return default + if isinstance(value, int): + return value + if isinstance(value, str): + try: + return int(value) + except ValueError: + return default + return default + + +def _safe_str(value: object, default: str) -> str: + return value if isinstance(value, str) else default + + +_ResponseModelT = TypeVar("_ResponseModelT", bound=BaseModel) + + +def _build_bag( + model_cls: type[_ResponseModelT], + **fields: object, # kwargs-ok: generic forwarder for extra-allow Responses payload models +) -> _ResponseModelT: + return model_cls(**fields) + + +@dataclass(frozen=True, slots=True) +class _ResponsesStreamItemState: + item_id: str + output_index: int + content_index: int = 0 + item_type: str = "message" + item_snapshot: str = "{}" + reasoning_summary: tuple[tuple[int, str], ...] = () + part_kind: str = "output_text" + accumulated_text: str = "" + output_item_added_seen: bool = False + content_part_added_seen: bool = False + leaf_done_seen: bool = False + content_part_done_seen: bool = False + output_item_done_seen: bool = False + + @property + def has_content_part(self) -> bool: + return self.item_type == "message" + + +_ItemStateMap = Mapping[int, _ResponsesStreamItemState] + + +class _ResponsesLifecycleGapFiller: + def __init__(self, *, model: str, response_id: str) -> None: + self._model = model + self._response_id = response_id + self._created_seen = False + self._in_progress_seen = False + self._items: _ItemStateMap = MappingProxyType({}) + + def expand(self, event: ResponsesAPIStreamingResponse) -> tuple[ResponsesAPIStreamingResponse, ...]: + ev = ResponsesAPIStreamEvents + etype = _obj_get(event, "type") + + if etype in (ev.RESPONSE_CREATED, ev.RESPONSE_IN_PROGRESS): + self._response_id = _safe_str(_obj_get(_obj_get(event, "response"), "id"), self._response_id) + self._created_seen = True + self._in_progress_seen = self._in_progress_seen or etype == ev.RESPONSE_IN_PROGRESS + return (event,) + if etype == ev.OUTPUT_ITEM_ADDED: + openers = self._response_openers() + self._observe_output_item_added(event) + return (*openers, event) + if etype == ev.CONTENT_PART_ADDED: + openers = self._response_openers() + self._observe_content_part_added(event) + return (*openers, event) + if etype in (ev.OUTPUT_TEXT_DELTA, ev.REFUSAL_DELTA): + openers = ( + *self._response_openers(), + *self._ensure_message_item(event, is_refusal=(etype == ev.REFUSAL_DELTA)), + ) + self._accumulate(event, _safe_str(_obj_get(event, "delta", ""), "")) + return (*openers, event) + if etype == ev.FUNCTION_CALL_ARGUMENTS_DELTA: + openers = ( + *self._response_openers(), + *self._ensure_function_call_item(event), + ) + self._accumulate(event, _safe_str(_obj_get(event, "delta", ""), "")) + return (*openers, event) + if etype in ( + "response.reasoning_summary_text.delta", + "response.reasoning_summary_text.done", + "response.reasoning_summary_part.added", + "response.reasoning_summary_part.done", + ): + self._observe_reasoning_summary(event, is_delta=etype == "response.reasoning_summary_text.delta") + return (event,) + if etype in ( + ev.OUTPUT_TEXT_DONE, + ev.REFUSAL_DONE, + ev.FUNCTION_CALL_ARGUMENTS_DONE, + ): + self._mark_seen(event, "leaf_done_seen") + return (event,) + if etype == ev.CONTENT_PART_DONE: + self._mark_seen(event, "content_part_done_seen") + return (event,) + if etype == ev.OUTPUT_ITEM_DONE: + self._observe_output_item_done(event) + return (event,) + if etype in (ev.RESPONSE_COMPLETED, ev.RESPONSE_INCOMPLETE, ev.RESPONSE_FAILED): + openers = self._response_openers() if self._items else () + return (*openers, *self._teardown(event), event) + return (event,) + + def _response_openers(self) -> tuple[BaseLiteLLMOpenAIResponseObject, ...]: + need_created = not self._created_seen + need_in_progress = not self._in_progress_seen + self._created_seen = True + self._in_progress_seen = True + return ( + *((self._status_event(is_created=True),) if need_created else ()), + *((self._status_event(is_created=False),) if need_in_progress else ()), + ) + + def _status_event(self, *, is_created: bool) -> BaseLiteLLMOpenAIResponseObject: + # Without an upstream opener, this ID is temporary; preserve the terminal ID for follow-up requests. + response = _build_bag( + ResponsesAPIResponse, + id=self._response_id, + created_at=int(time.time()), + model=self._model, + object="response", + status="in_progress", + output=(), + ) + if is_created: + return ResponseCreatedEvent(type=ResponsesAPIStreamEvents.RESPONSE_CREATED, response=response) + return ResponseInProgressEvent(type=ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, response=response) + + def _item_for(self, event: object) -> _ResponsesStreamItemState: + output_index: Final = _safe_int(_obj_get(event, "output_index", 0), 0) + existing: Final = self._items.get(output_index) + if existing is not None: + return existing + state: Final = _ResponsesStreamItemState( + item_id=_safe_str(_obj_get(event, "item_id", ""), "") or self._response_id, + output_index=output_index, + content_index=_safe_int(_obj_get(event, "content_index", 0), 0), + ) + return self._store_item(state) + + def _store_item(self, state: _ResponsesStreamItemState) -> _ResponsesStreamItemState: + self._items = MappingProxyType({**self._items, state.output_index: state}) + return state + + def _ensure_message_item(self, event: object, *, is_refusal: bool) -> tuple[BaseLiteLLMOpenAIResponseObject, ...]: + previous: Final = self._item_for(event) + need_item: Final = not previous.output_item_added_seen + need_part: Final = not previous.content_part_added_seen + state: Final = self._store_item( + replace( + previous, + item_type="message", + part_kind="refusal" if is_refusal else "output_text", + output_item_added_seen=True, + content_part_added_seen=True, + ) + ) + return ( + *((self._build_output_item_added(state),) if need_item else ()), + *((self._build_content_part_added(state),) if need_part else ()), + ) + + def _ensure_function_call_item(self, event: object) -> tuple[BaseLiteLLMOpenAIResponseObject, ...]: + previous: Final = self._item_for(event) + state: Final = self._store_item(replace(previous, item_type="function_call", output_item_added_seen=True)) + return () if previous.output_item_added_seen else (self._build_output_item_added(state),) + + def _accumulate(self, event: object, delta: str) -> None: + state: Final = self._item_for(event) + self._store_item(replace(state, accumulated_text=state.accumulated_text + delta)) + + def _mark_seen(self, event: object, flag: Literal["leaf_done_seen", "content_part_done_seen"]) -> None: + state: Final = self._item_for(event) + self._store_item( + replace( + state, + leaf_done_seen=state.leaf_done_seen or flag == "leaf_done_seen", + content_part_done_seen=state.content_part_done_seen or flag == "content_part_done_seen", + ) + ) + + def _observe_output_item_added(self, event: object) -> None: + output_index: Final = _safe_int(_obj_get(event, "output_index", 0), 0) + item: Final = _obj_get(event, "item") + item_type: Final = _obj_get(item, "type") + item_id: Final = ( + _safe_str(_obj_get(item, "id", ""), "") + or _safe_str(_obj_get(event, "item_id", ""), "") + or self._response_id + ) + state: Final = self._items.get(output_index) or _ResponsesStreamItemState( + item_id=item_id, output_index=output_index + ) + self._store_item( + replace( + state, + item_id=item_id, + item_type="message" if item_type == "refusal" else _safe_str(item_type, "message"), + item_snapshot=item.model_dump_json() if isinstance(item, BaseModel) else json.dumps(item), + output_item_added_seen=True, + ) + ) + + def _observe_content_part_added(self, event: object) -> None: + self._store_item(replace(self._item_for(event), content_part_added_seen=True)) + + def _observe_output_item_done(self, event: object) -> None: + output_index: Final = _safe_int(_obj_get(event, "output_index", 0), 0) + state: Final = self._items.get(output_index) + if state is not None: + self._store_item(replace(state, output_item_done_seen=True)) + + def _observe_reasoning_summary(self, event: object, *, is_delta: bool) -> None: + state: Final = self._items.get(_safe_int(_obj_get(event, "output_index"), 0)) + if state is None or state.item_type != "reasoning": + return + index: Final = _safe_int(_obj_get(event, "summary_index"), 0) + previous: Final = next((text for part_index, text in state.reasoning_summary if part_index == index), "") + text: Final = _safe_str( + _obj_get(event, "delta", _obj_get(event, "text", _obj_get(_obj_get(event, "part"), "text"))), "" + ) + self._store_item( + replace( + state, + reasoning_summary=( + *((part_index, value) for part_index, value in state.reasoning_summary if part_index != index), + (index, previous + text if is_delta else text), + ), + ) + ) + + def _teardown(self, terminal_event: object) -> tuple[BaseLiteLLMOpenAIResponseObject, ...]: + return tuple( + event for _, state in sorted(self._items.items()) for event in self._item_teardown(state, terminal_event) + ) + + def _item_teardown( + self, state: _ResponsesStreamItemState, terminal_event: object + ) -> tuple[BaseLiteLLMOpenAIResponseObject, ...]: + if state.output_item_done_seen: + return () + if state.item_type not in ("message", "function_call"): + self._store_item(replace(state, output_item_done_seen=True)) + return (self._build_other_item_done(state, terminal_event),) + need_leaf: Final = not state.leaf_done_seen + need_content_part: Final = state.has_content_part and not state.content_part_done_seen + self._store_item(replace(state, leaf_done_seen=True, content_part_done_seen=True, output_item_done_seen=True)) + return ( + *((self._build_leaf_done(state),) if need_leaf else ()), + *((self._build_content_part_done(state),) if need_content_part else ()), + self._build_output_item_done(state), + ) + + def _build_other_item_done(self, state: _ResponsesStreamItemState, terminal_event: object) -> OutputItemDoneEvent: + response: Final = _obj_get(terminal_event, "response") + terminal_item: Final = next( + ( + item + for item in _json_array_or_empty(_obj_get(response, "output")) + if _obj_get(item, "id") == state.item_id + ), + None, + ) + payload: Final = _load_json_value( + state.item_snapshot + if terminal_item is None + else terminal_item.model_dump_json() + if isinstance(terminal_item, BaseModel) + else json.dumps(terminal_item) + ) + fields: Final = MappingProxyType(payload) if _is_json_object(payload) else EMPTY_MAPPING + summary: Final = MappingProxyType( + { + "summary": tuple( + _build_bag(BaseLiteLLMOpenAIResponseObject, type="summary_text", text=text) + for _, text in sorted(state.reasoning_summary) + ) + } + if state.reasoning_summary + else {} + ) + return OutputItemDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=state.output_index, + item=_build_bag( + BaseLiteLLMOpenAIResponseObject, + **MappingProxyType( + {**fields, "id": state.item_id, "type": state.item_type, "status": "completed", **summary} + ), + ), + ) + + def _build_output_item_added(self, state: _ResponsesStreamItemState) -> OutputItemAddedEvent: + if state.has_content_part: + item = _build_bag( + BaseLiteLLMOpenAIResponseObject, + id=state.item_id, + type="message", + status="in_progress", + role="assistant", + content=(), + ) + else: + item = _build_bag( + BaseLiteLLMOpenAIResponseObject, + id=state.item_id, + type="function_call", + status="in_progress", + ) + return OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=state.output_index, + item=item, + ) + + def _build_content_part_added(self, state: _ResponsesStreamItemState) -> ContentPartAddedEvent: + if state.part_kind == "refusal": + part = _build_bag(BaseLiteLLMOpenAIResponseObject, type="refusal", refusal="") + else: + part = _build_bag( + BaseLiteLLMOpenAIResponseObject, + type="output_text", + text="", + annotations=(), + ) + return ContentPartAddedEvent( + type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED, + item_id=state.item_id, + output_index=state.output_index, + content_index=state.content_index, + part=part, + ) + + def _build_leaf_done(self, state: _ResponsesStreamItemState) -> BaseLiteLLMOpenAIResponseObject: + if not state.has_content_part: + return FunctionCallArgumentsDoneEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=state.item_id, + output_index=state.output_index, + arguments=state.accumulated_text, + ) + if state.part_kind == "refusal": + return RefusalDoneEvent( + type=ResponsesAPIStreamEvents.REFUSAL_DONE, + item_id=state.item_id, + output_index=state.output_index, + content_index=state.content_index, + refusal=state.accumulated_text, + ) + return OutputTextDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE, + item_id=state.item_id, + output_index=state.output_index, + content_index=state.content_index, + text=state.accumulated_text, + ) + + def _build_content_part_done(self, state: _ResponsesStreamItemState) -> ContentPartDoneEvent: + if state.part_kind == "refusal": + part: BaseLiteLLMOpenAIResponseObject = _build_bag( + ContentPartDonePartRefusal, + type="refusal", + refusal=state.accumulated_text, + ) + else: + part = _build_bag( + ContentPartDonePartOutputText, + type="output_text", + text=state.accumulated_text, + annotations=(), + logprobs=None, + ) + return ContentPartDoneEvent( + type=ResponsesAPIStreamEvents.CONTENT_PART_DONE, + item_id=state.item_id, + output_index=state.output_index, + content_index=state.content_index, + part=part, + ) + + def _build_output_item_done(self, state: _ResponsesStreamItemState) -> OutputItemDoneEvent: + if not state.has_content_part: + item = _build_bag( + BaseLiteLLMOpenAIResponseObject, + id=state.item_id, + type="function_call", + status="completed", + arguments=state.accumulated_text, + ) + else: + if state.part_kind == "refusal": + content_part = _build_bag( + BaseLiteLLMOpenAIResponseObject, + type="refusal", + refusal=state.accumulated_text, + ) + else: + content_part = _build_bag( + BaseLiteLLMOpenAIResponseObject, + type="output_text", + text=state.accumulated_text, + annotations=(), + ) + item = _build_bag( + BaseLiteLLMOpenAIResponseObject, + id=state.item_id, + type="message", + status="completed", + role="assistant", + content=(content_part,), + ) + return OutputItemDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=state.output_index, + item=item, + ) + + class BaseResponsesAPIStreamingIterator: """ Base class for streaming iterators that process responses from the Responses API. @@ -254,6 +708,12 @@ class BaseResponsesAPIStreamingIterator: self._persist_completed_response_before_logging = True self._stream_created_time: float = time.time() + self._pending_events: tuple[ResponsesAPIStreamingResponse, ...] = () + self._lifecycle_gap_filler = _ResponsesLifecycleGapFiller( + model=model or "", + response_id=f"resp_{uuid.uuid4().hex}", + ) + # track request context for hooks self.litellm_metadata = litellm_metadata self.custom_llm_provider = custom_llm_provider @@ -870,6 +1330,11 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): try: self._check_max_streaming_duration() while True: + if self._pending_events: + pending_event, self._pending_events = self._pending_events[0], self._pending_events[1:] + self._yielded_first_chunk = True + return pending_event + # Get the next chunk from the stream try: sse = await self.stream_iterator.__anext__() @@ -884,14 +1349,10 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopAsyncIteration elif result is not None: self._maybe_raise_for_error_event(result) - # Await hook directly instead of run_async_function - # (which spawns a thread + event loop per call) - result = await self._call_post_streaming_deployment_hook( - chunk=result, + # Accumulate post-hook deltas so teardown cannot restore redacted text. + self._pending_events = self._lifecycle_gap_filler.expand( + await self._call_post_streaming_deployment_hook(chunk=result) ) - self._yielded_first_chunk = True - return result - # If result is None, continue the loop to get the next chunk except StopAsyncIteration: # Normal end of stream - don't log as failure @@ -952,6 +1413,11 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): try: self._check_max_streaming_duration() while True: + if self._pending_events: + pending_event, self._pending_events = self._pending_events[0], self._pending_events[1:] + self._yielded_first_chunk = True + return pending_event + # Get the next chunk from the stream try: sse = next(self.stream_iterator) @@ -966,14 +1432,13 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator): raise StopIteration elif result is not None: self._maybe_raise_for_error_event(result) - # Sync path: use run_async_function for the hook - result = run_async_function( - async_function=self._call_post_streaming_deployment_hook, - chunk=result, + # Accumulate post-hook deltas so teardown cannot restore redacted text. + self._pending_events = self._lifecycle_gap_filler.expand( + run_async_function( + async_function=self._call_post_streaming_deployment_hook, + chunk=result, + ) ) - self._yielded_first_chunk = True - return result - # If result is None, continue the loop to get the next chunk except StopIteration: # Normal end of stream - don't log as failure diff --git a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py index bd617587cf3..f3645347c91 100644 --- a/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py +++ b/tests/llm_responses_api_testing/test_base_responses_api_streaming_iterator.py @@ -2,12 +2,12 @@ Unit tests for BaseResponsesAPIStreamingIterator Tests core functionality including: -1. Processing chunks and handling ResponseCompletedEvent +1. Processing chunks and handling ResponseCompletedEvent 2. Ensuring _update_responses_api_response_id_with_model_id is called for final chunk 3. Verifying ID update is NOT called for non-final chunks (delta events) 4. Edge case handling for invalid JSON, empty chunks, and [DONE] markers -These tests ensure the streaming iterator correctly processes response chunks +These tests ensure the streaming iterator correctly processes response chunks and applies model ID updates only to completed responses, as required for proper response tracking and logging. """ @@ -429,8 +429,8 @@ class TestBaseResponsesAPIStreamingIterator: except StopAsyncIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + assert mock_delta_event in chunks_received + assert chunks_received[-1] is mock_delta_event # CRITICAL: Verify that failure handlers were NOT called # StopAsyncIteration is a normal end of stream, not a failure @@ -490,8 +490,8 @@ class TestBaseResponsesAPIStreamingIterator: except StopIteration: pass # This is expected - # Verify we got the chunk - assert len(chunks_received) == 1 + assert mock_delta_event in chunks_received + assert chunks_received[-1] is mock_delta_event # CRITICAL: Verify that failure handlers were NOT called # StopIteration is a normal end of stream, not a failure diff --git a/tests/test_litellm/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py b/tests/test_litellm/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py index b9408d44e9a..59d47a77597 100644 --- a/tests/test_litellm/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/responses/test_fireworks_ai_responses_transformation.py @@ -450,6 +450,34 @@ def test_streaming_responses_call_hits_native_endpoint_and_yields_every_firework True, True, ) - assert tuple(event.type for event in received) == tuple(event["type"] for event in FIREWORKS_SSE_EVENTS) + assert tuple(event.type for event in received) == ( + "response.created", + "response.in_progress", + "response.output_item.added", + "response.reasoning_summary_text.delta", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.delta", + "response.output_item.done", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.completed", + ) + upstream_events: Final = tuple(received[index] for index in (0, 2, 3, 4, 6, 7, 12)) + assert tuple(event.type for event in upstream_events) == tuple(event["type"] for event in FIREWORKS_SSE_EVENTS) + assert tuple(event.sequence_number for event in upstream_events) == tuple(range(len(FIREWORKS_SSE_EVENTS))) + assert received[0].response.id == received[1].response.id == received[-1].response.id + assert received[2].item.type == "reasoning" + assert received[3].delta == "pong" + assert received[8].item.type == "reasoning" + assert received[8].item.id == "rs_1" + assert json.loads(received[8].model_dump_json())["item"]["summary"][0]["text"] == "pong" + assert received[9].text == "pong" + assert received[10].part.text == "pong" + assert received[11].item.type == "message" + assert received[11].item.id == "msg_1" + assert received[11].output_index == 1 assert "".join(event.delta for event in received if event.type == "response.output_text.delta") == "pong" assert received[-1].response.usage.output_tokens == 89 diff --git a/tests/test_litellm/responses/test_streaming_iterator.py b/tests/test_litellm/responses/test_streaming_iterator.py index c226c0b4d09..d248d65480e 100644 --- a/tests/test_litellm/responses/test_streaming_iterator.py +++ b/tests/test_litellm/responses/test_streaming_iterator.py @@ -4,24 +4,38 @@ completion_start_time on the first chunk so downstream TTFT consumers completion_start_time = end_time.""" import json +from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Optional +from types import MappingProxyType +from typing import Final, Optional from unittest.mock import Mock, patch import httpx import pytest +import litellm from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.responses.streaming_iterator import ( ResponsesAPIStreamingIterator, SyncResponsesAPIStreamingIterator, + _obj_get, + _ResponsesLifecycleGapFiller, + _safe_int, + _safe_str, ) from litellm.types.llms.openai import ( + BaseLiteLLMOpenAIResponseObject, ResponseCompletedEvent, ResponsesAPIResponse, ResponsesAPIStreamEvents, + ResponsesAPIStreamingResponse, ) +from litellm.types.utils import CallTypes + +EV = ResponsesAPIStreamEvents +E = EV def _sse_event(payload: dict) -> bytes: @@ -537,3 +551,549 @@ async def test_streaming_logging_copy_fallback_leaves_caller_event_untouched(): assert logged == [iterator.completed_response] assert iterator.completed_response.response._hidden_params == {} + + +def _response_body(status: str, *, model: str = "gpt-5") -> Mapping[str, object]: + return MappingProxyType( + { + "id": "resp_real_upstream", + "object": "response", + "created_at": 1700000000, + "status": status, + "model": model, + "output": ( + MappingProxyType( + { + "id": "msg_1", + "type": "message", + "status": "completed", + "role": "assistant", + "content": ( + MappingProxyType({"type": "output_text", "text": "Hello world", "annotations": ()}), + ), + } + ), + ), + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": (), + } + ) + + +def _sse_frames(events: Sequence[Mapping[str, object]]) -> tuple[bytes, ...]: + return (*(f"data: {json.dumps(evt, default=dict)}\n\n".encode("utf-8") for evt in events), b"data: [DONE]\n\n") + + +def _make_logging_obj() -> Mock: + logging_obj: Final = Mock(spec=LiteLLMLoggingObj) + logging_obj.model_call_details = MappingProxyType({"litellm_params": MappingProxyType({})}) + logging_obj.completion_start_time = None + return logging_obj + + +def _iterator( + events: Sequence[Mapping[str, object]], *, sync: bool, model: str = "gpt-5" +) -> ResponsesAPIStreamingIterator | SyncResponsesAPIStreamingIterator: + response: Final = httpx.Response(200, content=b"".join(_sse_frames(events))) + cls: Final = SyncResponsesAPIStreamingIterator if sync else ResponsesAPIStreamingIterator + return cls( + response=response, + model=model, + responses_api_provider_config=OpenAIResponsesAPIConfig(), + logging_obj=_make_logging_obj(), + custom_llm_provider="openai", + ) + + +async def _drive( + events: Sequence[Mapping[str, object]], *, sync: bool, model: str = "gpt-5" +) -> tuple[ResponsesAPIStreamingResponse, ...]: + iterator: Final = _iterator(events, sync=sync, model=model) + if isinstance(iterator, SyncResponsesAPIStreamingIterator): + return tuple(iterator) + return tuple([chunk async for chunk in iterator]) + + +def _types(events: Sequence[ResponsesAPIStreamingResponse]) -> tuple[str, ...]: + return tuple(_safe_str(_obj_get(event, "type"), "") for event in events) + + +_TRUNCATED_TEXT_EVENTS: Final[tuple[Mapping[str, object], ...]] = ( + MappingProxyType( + { + "type": "response.output_text.delta", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "delta": "Hello", + } + ), + MappingProxyType( + { + "type": "response.output_text.delta", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "delta": " world", + } + ), + MappingProxyType({"type": "response.completed", "response": _response_body("completed")}), +) +_FULL_TEXT_EVENTS: Final[tuple[Mapping[str, object], ...]] = ( + MappingProxyType({"type": "response.created", "response": _response_body("in_progress")}), + MappingProxyType({"type": "response.in_progress", "response": _response_body("in_progress")}), + MappingProxyType( + { + "type": "response.output_item.added", + "output_index": 0, + "item": MappingProxyType( + {"id": "msg_1", "type": "message", "status": "in_progress", "role": "assistant", "content": ()} + ), + } + ), + MappingProxyType( + { + "type": "response.content_part.added", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "part": MappingProxyType({"type": "output_text", "text": "", "annotations": ()}), + } + ), + MappingProxyType( + { + "type": "response.output_text.delta", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "delta": "Hello world", + } + ), + MappingProxyType( + { + "type": "response.output_text.done", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "text": "Hello world", + } + ), + MappingProxyType( + { + "type": "response.content_part.done", + "item_id": "msg_1", + "output_index": 0, + "content_index": 0, + "part": MappingProxyType({"type": "output_text", "text": "Hello world", "annotations": ()}), + } + ), + MappingProxyType( + { + "type": "response.output_item.done", + "output_index": 0, + "item": MappingProxyType( + { + "id": "msg_1", + "type": "message", + "status": "completed", + "role": "assistant", + "content": (MappingProxyType({"type": "output_text", "text": "Hello world", "annotations": ()}),), + } + ), + } + ), + MappingProxyType({"type": "response.completed", "response": _response_body("completed")}), +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +async def test_truncated_text_stream_synthesizes_full_lifecycle(sync: bool) -> None: + collected: Final = await _drive(_TRUNCATED_TEXT_EVENTS, sync=sync) + types: Final = _types(collected) + assert types == ( + E.RESPONSE_CREATED, + E.RESPONSE_IN_PROGRESS, + E.OUTPUT_ITEM_ADDED, + E.CONTENT_PART_ADDED, + E.OUTPUT_TEXT_DELTA, + E.OUTPUT_TEXT_DELTA, + E.OUTPUT_TEXT_DONE, + E.CONTENT_PART_DONE, + E.OUTPUT_ITEM_DONE, + E.RESPONSE_COMPLETED, + ), types + output_item_added: Final = collected[2] + content_part_added: Final = collected[3] + assert output_item_added.item.id == "msg_1" + assert content_part_added.item_id == "msg_1" + assert content_part_added.output_index == 0 + assert content_part_added.content_index == 0 + output_text_done: Final = collected[6] + assert output_text_done.text == "Hello world" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +async def test_complete_stream_passes_through_without_duplication(sync: bool) -> None: + collected: Final = await _drive(_FULL_TEXT_EVENTS, sync=sync) + types: Final = _types(collected) + assert types == tuple((evt["type"] for evt in _FULL_TEXT_EVENTS)), types + assert len(collected) == len(_FULL_TEXT_EVENTS) + assert types.count(E.RESPONSE_CREATED) == 1 + assert types.count(E.OUTPUT_ITEM_ADDED) == 1 + assert types.count(E.CONTENT_PART_ADDED) == 1 + assert types.count(E.OUTPUT_ITEM_DONE) == 1 + + +@pytest.mark.asyncio +async def test_truncated_function_call_stream_synthesizes_item_lifecycle() -> None: + events: Final = ( + MappingProxyType( + { + "type": "response.function_call_arguments.delta", + "item_id": "fc_1", + "output_index": 0, + "delta": '{"city":', + } + ), + MappingProxyType( + {"type": "response.function_call_arguments.delta", "item_id": "fc_1", "output_index": 0, "delta": '"NYC"}'} + ), + MappingProxyType({"type": "response.completed", "response": _response_body("completed")}), + ) + collected: Final = await _drive(events, sync=False) + types: Final = _types(collected) + assert types == ( + E.RESPONSE_CREATED, + E.RESPONSE_IN_PROGRESS, + E.OUTPUT_ITEM_ADDED, + E.FUNCTION_CALL_ARGUMENTS_DELTA, + E.FUNCTION_CALL_ARGUMENTS_DELTA, + E.FUNCTION_CALL_ARGUMENTS_DONE, + E.OUTPUT_ITEM_DONE, + E.RESPONSE_COMPLETED, + ), types + assert E.CONTENT_PART_ADDED not in types + assert E.CONTENT_PART_DONE not in types + output_item_added: Final = collected[2] + assert output_item_added.item.type == "function_call" + assert output_item_added.item.id == "fc_1" + args_done: Final = collected[5] + assert args_done.arguments == '{"city":"NYC"}' + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +@pytest.mark.parametrize("complete_reasoning", (False, True), ids=("truncated-reasoning", "complete-reasoning")) +async def test_gpt_5_6_reasoning_stream_preserves_item_lifecycle(sync: bool, complete_reasoning: bool) -> None: + reasoning_events: Final = ( + MappingProxyType( + { + "type": "response.output_item.added", + "output_index": 0, + "item": MappingProxyType({"id": "rs_1", "type": "reasoning", "summary": ()}), + } + ), + MappingProxyType( + { + "type": "response.reasoning_summary_text.delta", + "item_id": "rs_1", + "output_index": 0, + "summary_index": 0, + "delta": "Thinking", + } + ), + MappingProxyType( + { + "type": "response.reasoning_summary_text.done", + "item_id": "rs_1", + "output_index": 0, + "summary_index": 0, + "sequence_number": 4, + "text": "Thinking", + } + ), + MappingProxyType( + { + "type": "response.output_item.done", + "output_index": 0, + "item": MappingProxyType( + { + "id": "rs_1", + "type": "reasoning", + "summary": (MappingProxyType({"type": "summary_text", "text": "Thinking"}),), + } + ), + } + ), + ) + message_events: Final = tuple( + ( + MappingProxyType( + { + **event, + **(MappingProxyType({"output_index": 1}) if "output_index" in event else MappingProxyType({})), + } + ) + for event in _FULL_TEXT_EVENTS[2:] + ) + ) + events: Final = tuple( + ( + MappingProxyType( + { + **event, + **( + MappingProxyType( + { + "response": _response_body( + "completed" if event["type"] == E.RESPONSE_COMPLETED else "in_progress", + model="gpt-5.6", + ) + } + ) + if "response" in event + else MappingProxyType({}) + ), + } + ) + for event in ( + *_FULL_TEXT_EVENTS[:2], + *(reasoning_events if complete_reasoning else reasoning_events[:2]), + *message_events, + ) + ) + ) + collected: Final = await _drive(events, sync=sync, model="gpt-5.6") + expected_types: Final = tuple((event["type"] for event in events)) + assert _types(collected) == ( + expected_types if complete_reasoning else (*expected_types[:-1], E.OUTPUT_ITEM_DONE, expected_types[-1]) + ) + assert collected[2].item.id == "rs_1" + assert collected[2].item.type == "reasoning" + assert collected[3].delta == "Thinking" + if complete_reasoning: + assert collected[4].text == "Thinking" + assert collected[5].item.type == "reasoning" + message_start: Final = 6 if complete_reasoning else 4 + assert collected[message_start].output_index == 1 + assert collected[message_start].item.id == "msg_1" + assert collected[message_start + 2].delta == "Hello world" + assert collected[message_start + 3].text == "Hello world" + assert collected[-1].response.model == "gpt-5.6" + assert E.FUNCTION_CALL_ARGUMENTS_DONE not in _types(collected) + reasoning_done: Final = next( + (event for event in collected if event.type == E.OUTPUT_ITEM_DONE and event.item.id == "rs_1") + ) + assert reasoning_done.item.type == "reasoning" + assert json.loads(reasoning_done.model_dump_json())["item"]["summary"][0]["text"] == "Thinking" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +@pytest.mark.parametrize("terminal_has_item", (False, True)) +@pytest.mark.parametrize("has_summary_deltas", (False, True)) +async def test_reasoning_teardown_preserves_summary_indices_and_encrypted_content( + sync: bool, terminal_has_item: bool, has_summary_deltas: bool +) -> None: + opening_item: Final = MappingProxyType( + {"id": "rs_1", "type": "reasoning", "summary": (), "encrypted_content": "opening-encrypted"} + ) + terminal_item: Final = MappingProxyType( + {**opening_item, "status": "completed", "encrypted_content": "final-encrypted"} + ) + events: Final = ( + *_FULL_TEXT_EVENTS[:2], + MappingProxyType({"type": E.OUTPUT_ITEM_ADDED, "output_index": 0, "item": opening_item}), + *( + MappingProxyType( + { + "type": "response.reasoning_summary_text.delta", + "item_id": "rs_1", + "output_index": 0, + "summary_index": index, + "delta": text, + } + ) + for index, text in ((1, "Second"), (0, "Think"), (0, "ing")) + if has_summary_deltas + ), + MappingProxyType( + { + "type": E.RESPONSE_COMPLETED, + "response": MappingProxyType( + {**_response_body("completed"), "output": (terminal_item,) if terminal_has_item else ()} + ), + } + ), + ) + collected: Final = await _drive(events, sync=sync) + done: Final = collected[-2] + wire: Final = json.loads(done.model_dump_json(exclude_none=True, exclude_unset=True)) + + assert done.type == E.OUTPUT_ITEM_DONE + assert done.output_index == 0 + assert wire["item"]["id"] == "rs_1" + assert wire["item"]["type"] == "reasoning" + assert wire["item"]["status"] == "completed" + assert wire["item"]["encrypted_content"] == ("final-encrypted" if terminal_has_item else "opening-encrypted") + assert tuple(part["text"] for part in wire["item"]["summary"]) == ( + ("Thinking", "Second") if has_summary_deltas else () + ) + assert _types(collected).count(E.OUTPUT_ITEM_DONE) == 1 + assert E.FUNCTION_CALL_ARGUMENTS_DONE not in _types(collected) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +async def test_stream_without_item_events_preserves_response_status_events(sync: bool) -> None: + events: Final = (_FULL_TEXT_EVENTS[0], _FULL_TEXT_EVENTS[-1]) + collected: Final = await _drive(events, sync=sync) + assert _types(collected) == (E.RESPONSE_CREATED, E.RESPONSE_COMPLETED) + assert collected[0].response.id == collected[1].response.id + + +@pytest.mark.asyncio +async def test_synthesized_events_survive_proxy_serialization() -> None: + collected: Final = await _drive(_TRUNCATED_TEXT_EVENTS, sync=False) + required_by_type: Final = MappingProxyType( + { + E.OUTPUT_ITEM_ADDED: ("type", "output_index", "item"), + E.CONTENT_PART_ADDED: ("type", "item_id", "output_index", "content_index", "part"), + E.OUTPUT_TEXT_DONE: ("type", "item_id", "output_index", "content_index", "text"), + E.CONTENT_PART_DONE: ("type", "item_id", "output_index", "content_index", "part"), + E.OUTPUT_ITEM_DONE: ("type", "output_index", "item"), + } + ) + seen_types: Final = frozenset((event.type for event in collected if event.type in required_by_type)) + serialized: Final = tuple( + (event.type, json.loads(event.model_dump_json(exclude_none=True, exclude_unset=True))) + for event in collected + if event.type in required_by_type + ) + assert all(field in wire for event_type, wire in serialized for field in required_by_type[event_type]) + assert seen_types == frozenset(required_by_type) + + +class _RedactingDeploymentHook: + REDACTION: Final = "[REDACTED]" + + async def async_post_call_streaming_deployment_hook( + self, + *, + request_data: Mapping[str, object], + response_chunk: ResponsesAPIStreamingResponse, + call_type: CallTypes | None, + ) -> ResponsesAPIStreamingResponse: + if getattr(response_chunk, "type", None) == E.OUTPUT_TEXT_DELTA: + return response_chunk.model_copy(update=MappingProxyType({"delta": self.REDACTION})) + return response_chunk + + +@pytest.fixture +def redacting_deployment_hook(monkeypatch: pytest.MonkeyPatch) -> _RedactingDeploymentHook: + hook: Final = _RedactingDeploymentHook() + monkeypatch.setattr(litellm, "callbacks", (*litellm.callbacks, hook)) + return hook + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +async def test_streaming_hook_governs_synthesized_teardown( + sync: bool, redacting_deployment_hook: _RedactingDeploymentHook +) -> None: + redacted: Final = _RedactingDeploymentHook.REDACTION * 2 + collected: Final = await _drive(_TRUNCATED_TEXT_EVENTS, sync=sync) + by_type: Final = MappingProxyType( + { + event_type: tuple((event for event in collected if event.type == event_type)) + for event_type in _types(collected) + } + ) + assert tuple((d.delta for d in by_type[E.OUTPUT_TEXT_DELTA])) == ( + _RedactingDeploymentHook.REDACTION, + _RedactingDeploymentHook.REDACTION, + ) + assert by_type[E.OUTPUT_TEXT_DONE][0].text == redacted + assert by_type[E.CONTENT_PART_DONE][0].part.text == redacted + assert by_type[E.OUTPUT_ITEM_DONE][0].item.content[0].text == redacted + assert all((getattr(e, "text", None) != "Hello world" for e in collected)) + + +_TRUNCATED_REFUSAL_EVENTS: Final[tuple[Mapping[str, object], ...]] = ( + MappingProxyType( + {"type": "response.refusal.delta", "item_id": "msg_r", "output_index": 0, "content_index": 0, "delta": "I can"} + ), + MappingProxyType( + { + "type": "response.refusal.delta", + "item_id": "msg_r", + "output_index": 0, + "content_index": 0, + "delta": "not help", + } + ), + MappingProxyType({"type": "response.completed", "response": _response_body("completed")}), +) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sync", (False, True), ids=("async", "sync")) +async def test_truncated_refusal_stream_synthesizes_lifecycle(sync: bool) -> None: + collected: Final = await _drive(_TRUNCATED_REFUSAL_EVENTS, sync=sync) + types: Final = _types(collected) + assert types == ( + E.RESPONSE_CREATED, + E.RESPONSE_IN_PROGRESS, + E.OUTPUT_ITEM_ADDED, + E.CONTENT_PART_ADDED, + E.REFUSAL_DELTA, + E.REFUSAL_DELTA, + E.REFUSAL_DONE, + E.CONTENT_PART_DONE, + E.OUTPUT_ITEM_DONE, + E.RESPONSE_COMPLETED, + ), types + assert collected[3].part.type == "refusal" + assert collected[6].refusal == "I cannot help" + assert collected[7].part.refusal == "I cannot help" + assert collected[8].item.content[0].refusal == "I cannot help" + + +def test_obj_get_handles_dict_object_and_none() -> None: + assert _obj_get(MappingProxyType({"a": 1}), "a") == 1 + assert _obj_get(MappingProxyType({"a": 1}), "missing", "d") == "d" + assert _obj_get(None, "a", "d") == "d" + + class _Obj: + x: Final = 5 + + assert _obj_get(_Obj(), "x") == 5 + assert _obj_get(_Obj(), "y", "fallback") == "fallback" + + +def test_safe_int_narrows_dynamic_values() -> None: + assert _safe_int(3, 0) == 3 + assert _safe_int(True, 9) == 9 + assert _safe_int("5", 0) == 5 + assert _safe_int("nope", 7) == 7 + assert _safe_int(1.5, 4) == 4 + + +@pytest.mark.parametrize("event_type", ("response.some_unhandled_event", "response.reasoning_summary_text.delta")) +def test_gap_filler_passes_events_without_items_through(event_type: str) -> None: + gap_filler: Final = _ResponsesLifecycleGapFiller(model="m", response_id="resp_x") + event: Final = BaseLiteLLMOpenAIResponseObject.model_validate( + MappingProxyType( + { + "type": event_type, + "item_id": "rs_orphan", + "output_index": 0, + "summary_index": 0, + "delta": "Thinking", + } + ) + ) + assert gap_filler.expand(event) == (event,) diff --git a/tests/test_litellm/responses/test_streaming_iterator_error_events.py b/tests/test_litellm/responses/test_streaming_iterator_error_events.py index ad74861c096..aa459584830 100644 --- a/tests/test_litellm/responses/test_streaming_iterator_error_events.py +++ b/tests/test_litellm/responses/test_streaming_iterator_error_events.py @@ -212,7 +212,14 @@ async def test_async_iterator_error_after_first_chunk_carries_generated_content( with pytest.raises(MidStreamFallbackError) as exc_info: await _drain() - assert len(chunks) == 2 + assert [chunk.type for chunk in chunks] == [ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.delta", + ] assert exc_info.value.status_code == 500 assert exc_info.value.is_pre_first_chunk is False assert exc_info.value.generated_content == "hello world"