This commit is contained in:
Ali Khan 2026-09-08 17:34:03 +00:00 • committed by GitHub
commit 198921ac3b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 1084 additions and 24 deletions

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,)

View file

@ -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"