mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 40fef3cedb into 9b74e6f34e
This commit is contained in:
commit
198921ac3b
5 changed files with 1084 additions and 24 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue