fix(responses-api): raise APIError on in-stream error events; widen ErrorEventError.param to accept dict (#32835)

* fix(responses-api): raise APIError on in-stream error events; widen ErrorEventError.param

- BaseResponsesAPIStreamingIterator._maybe_raise_for_error_event inspects each
  chunk and raises litellm.APIError for type=error and type=response.failed events
  so callers see an exception instead of a benign stream chunk
- rate_limit* codes map to 429; client error codes (invalid_request_error,
  context_length_exceeded, etc.) map to 400; all other codes default to 500;
  raw integer codes are never used as-is as HTTP status codes
- ErrorEventError.param widened from Optional[str] to Optional[Union[str, Dict]]
  to prevent Pydantic ValidationError on dict-typed param payloads silently
  dropping error events before any type inspection

* test(responses-api): add streaming iterator error event tests to CI-covered path

* test(responses-api): cover response.failed, dict-error, null-error, and sync iterator paths

* test(responses-api): set completion_start_time on mock logging objects for internal staging _process_chunk

* fix(responses-api): map insufficient_quota to 429, derive failed-response log status from error code, and record failed-stream usage for spend accounting

insufficient_quota moves out of the 400 bucket; OpenAI returns HTTP 429 for it and the non-streaming exception mapping treats 429 as RateLimitError, so the in-stream mapping now agrees

_handle_logging_failed_response previously hardcoded APIError(status_code=500), so a rate-limited response.failed was logged to integrations as 500 while the caller saw 429; it now shares the same error-code-to-status mapping via _error_event_fields and _status_code_for_error_code

usage carried on a response.failed event is now stashed as combined_usage_object with its computed cost on the logging object before failure handlers run, reusing the mid-stream-interruption spend recovery path (_failure_handler_helper_fn, proxy post_call_failure_hook, _ProxyDBLogger), so failed streams count their billed tokens instead of logging zero cost

dedupe: TestMaybeRaiseForErrorEvent in tests/llm_responses_api_testing duplicated tests/test_litellm/responses/test_streaming_iterator_error_events.py, which is the canonical mirrored location and CI-covered via test-unit-responses-caching-types; the duplicate class is removed

* fix(responses-api): wrap retriable in-stream errors in MidStreamFallbackError and map error type field to status

Mirror chat streaming semantics from _handle_stream_fallback_error: 429 and
5xx in-stream error events now raise MidStreamFallbackError carrying the
mapped APIError so the router's FallbackResponsesStreamWrapper triggers
mid-stream fallback and cooldown; non-retriable 4xx still raise APIError
directly. Status mapping now reads both the OpenAI error type and code
fields, so type-classified client errors (e.g. invalid_request_error with
code invalid_prompt) map to 400 instead of falling through to 500.

* fix(responses-api): accumulate streamed output text so mid-stream fallback continues instead of restarting

MidStreamFallbackError was always raised with generated_content="", so the
router's stream_with_fallbacks treated every mid-stream error as pre-first-chunk
and retried with the original input, streaming duplicated content to clients
that had already received partial output. The iterators now accumulate
response.output_text.delta text (mirroring chat's response_uptil_now) and pass
it as generated_content, letting the router build a continuation input via
_build_responses_continuation_input.

* test(responses-api): pin in-stream token limit error to raised APIError

---------

Co-authored-by: Deepanshu <deepanshu.lulla@alpha-sense.com>
This commit is contained in:
Mateo Wang 2026-07-10 20:25:38 -07:00 committed by GitHub
parent 5e23a5ab05
commit 249a999506
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 691 additions and 13 deletions

View file

@ -17,6 +17,7 @@ from litellm.constants import (
LITELLM_MAX_STREAMING_DURATION_SECONDS,
STREAM_SSE_DONE_STRING,
)
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
@ -26,7 +27,7 @@ from litellm.litellm_core_utils.llm_response_utils.response_metadata import (
)
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.responses.utils import ResponseAPILoggingUtils, ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIStreamEvents
from litellm.types.utils import CallTypes
from litellm.utils import async_post_call_success_deployment_hook
@ -47,6 +48,44 @@ def _log_background_task_failure(task: "asyncio.Task[Any]", *, task_name: str) -
verbose_logger.error("%s failed: %s", task_name, exception)
_CLIENT_ERROR_CODES: frozenset[str] = frozenset(
(
"invalid_request_error",
"context_length_exceeded",
"content_policy_violation",
"model_not_found",
)
)
def _error_event_fields(error_obj: object) -> tuple[str, Optional[str], Optional[str]]:
if isinstance(error_obj, dict):
raw_message = error_obj.get("message")
raw_type = error_obj.get("type")
raw_code = error_obj.get("code")
elif error_obj is not None:
raw_message = getattr(error_obj, "message", None)
raw_type = getattr(error_obj, "type", None)
raw_code = getattr(error_obj, "code", None)
else:
raw_message = None
raw_type = None
raw_code = None
message = str(raw_message) if raw_message is not None else "Response API in-stream error"
error_type = raw_type if isinstance(raw_type, str) else None
code = raw_code if isinstance(raw_code, str) else None
return message, error_type, code
def _status_code_for_error_fields(error_type: Optional[str], error_code: Optional[str]) -> int:
fields = tuple(field for field in (error_type, error_code) if field is not None)
if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields):
return 429
if any(field in _CLIENT_ERROR_CODES for field in fields):
return 400
return 500
class BaseResponsesAPIStreamingIterator:
"""
Base class for streaming iterators that process responses from the Responses API.
@ -73,6 +112,8 @@ class BaseResponsesAPIStreamingIterator:
self.completed_response: Optional[Any] = None
self.start_time = getattr(logging_obj, "start_time", datetime.now())
self._failure_handled = False # Track if failure handler has been called
self._yielded_first_chunk = False
self._generated_content = ""
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit: Optional[bool] = None
@ -160,6 +201,10 @@ class BaseResponsesAPIStreamingIterator:
# Encode container_id on streaming events so proxy/UI follow-ups route correctly
_event_type = getattr(openai_responses_api_chunk, "type", None)
if _event_type == ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA:
_delta = getattr(openai_responses_api_chunk, "delta", None)
if isinstance(_delta, str):
self._generated_content += _delta
_stream_model_id = (
self.litellm_metadata.get("model_info", {}).get("id") if self.litellm_metadata else None
)
@ -327,17 +372,66 @@ class BaseResponsesAPIStreamingIterator:
"""
response_obj = getattr(self.completed_response, "response", None) if self.completed_response else None
error_info = getattr(response_obj, "error", None) if response_obj else None
error_message = "Response failed"
if isinstance(error_info, dict):
error_message = error_info.get("message", str(error_info))
error_message, error_type, error_code = _error_event_fields(error_info)
self._record_failed_response_usage(response_obj)
exception = litellm.APIError(
status_code=500,
status_code=_status_code_for_error_fields(error_type, error_code),
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
)
self._handle_failure(exception)
def _record_failed_response_usage(self, response_obj: Optional[Any]) -> None:
if response_obj is None or self.logging_obj is None:
return
usage_obj = getattr(response_obj, "usage", None)
if usage_obj is None:
return
try:
self.logging_obj.model_call_details["combined_usage_object"] = (
ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(usage_obj)
)
except (TypeError, ValueError) as usage_error:
verbose_logger.debug(
"could not record usage for failed responses stream: %s",
usage_error,
)
return
self.logging_obj.model_call_details["response_cost"] = (
self.logging_obj._response_cost_calculator(result=response_obj) or 0.0
)
def _maybe_raise_for_error_event(self, result: object) -> None:
chunk_type = getattr(result, "type", None)
if chunk_type not in ("error", "response.failed"):
return
error_obj: object = (
getattr(getattr(result, "response", None), "error", None)
if chunk_type == "response.failed"
else getattr(result, "error", None)
)
error_message, error_type, error_code = _error_event_fields(error_obj)
status_code = _status_code_for_error_fields(error_type, error_code)
mapped_exception = litellm.APIError(
status_code=status_code,
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
)
if 400 <= status_code < 500 and status_code != 429:
raise mapped_exception
raise MidStreamFallbackError(
message=str(mapped_exception),
model=self.model or "",
llm_provider=self.custom_llm_provider or "",
original_exception=mapped_exception,
generated_content=self._generated_content,
is_pre_first_chunk=not self._yielded_first_chunk,
)
def _get_completed_response_object(self) -> Optional[Any]:
openai_types = _get_openai_response_types()
completed_response = self.completed_response
@ -611,11 +705,13 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
if self.finished:
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,
)
self._yielded_first_chunk = True
return result
# If result is None, continue the loop to get the next chunk
@ -685,11 +781,13 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
if self.finished:
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,
)
self._yielded_first_chunk = True
return result
# If result is None, continue the loop to get the next chunk

View file

@ -2344,6 +2344,8 @@ class Router:
self.completed_response = None
self.start_time = getattr(source_iterator, "start_time", datetime.now())
self._failure_handled = False
self._yielded_first_chunk = False
self._generated_content = ""
self._completed_response_cached = False
self._completed_response_logged = False
self._completed_response_cache_hit = None

View file

@ -1720,7 +1720,7 @@ class ErrorEventError(BaseLiteLLMOpenAIResponseObject):
type: str # e.g., 'invalid_request_error'
code: str # e.g., 'context_length_exceeded'
message: str
param: Optional[str] = None
param: Optional[Union[str, Dict[str, Any]]] = None
class ErrorEvent(BaseLiteLLMOpenAIResponseObject):

View file

@ -1628,23 +1628,27 @@ async def test_openai_responses_api_token_limit_error():
"""
Relevant issue: https://github.com/BerriAI/litellm/issues/15785
When this fails you'll see:
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent"
in the console.
Parsing the in-stream ErrorEvent must not raise
"pydantic_core._pydantic_core.ValidationError: 3 validation errors for ErrorEvent".
The iterator now surfaces the event as litellm.APIError with status 400
(invalid_request_error is a non-retriable client error, so no
MidStreamFallbackError wrapping) carrying the provider's message.
"""
litellm._turn_on_debug()
# Generate text with >400k tokens to trigger token limit error
oversized_text = "This is a test sentence. " * 50000 # ~400k tokens
# This will raise ValidationError instead of showing the real error
response = await litellm.aresponses(
model="gpt-5-mini", input=oversized_text, stream=True
)
async for event in response:
print(event) # Never reaches here - ValidationError is raised
with pytest.raises(litellm.APIError) as exc_info:
async for event in response:
print(event)
assert exc_info.value.status_code == 400
assert "exceeds the context window" in str(exc_info.value)
async def test_openai_streaming_logging():

View file

@ -266,3 +266,220 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator():
)
assert out is wrapped
mock_wrap.assert_awaited_once()
@pytest.mark.asyncio
async def test_aresponses_fallback_on_in_stream_error_event():
"""A retriable in-stream error event (429) must trigger the router's mid-stream
fallback path: the wrapper catches MidStreamFallbackError raised by the source
iterator and yields the fallback stream instead of surfacing the error."""
import json
from unittest.mock import Mock
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.types.llms.openai import ErrorEvent, ErrorEventError
router = _make_router()
error_payload = {
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"},
}
sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode()
async def mock_aiter_bytes():
yield sse_bytes
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_config.transform_streaming_response.return_value = ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(type="tokens", code="rate_limit_exceeded", message="rate limited"),
)
source = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
fallback_event = _make_completed_event(1, 1, 2)
class _FallbackStream:
def __init__(self) -> None:
self._done = False
def __aiter__(self):
return self
async def __anext__(self):
if self._done:
raise StopAsyncIteration
self._done = True
return fallback_event
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(return_value=_FallbackStream()),
) as mock_fallback:
wrapped = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={"model": "primary", "input": "original question"},
)
collected = [ev async for ev in wrapped]
assert collected == [fallback_event]
mock_fallback.assert_awaited_once()
raised = mock_fallback.await_args.kwargs["e"]
assert isinstance(raised, MidStreamFallbackError)
assert raised.status_code == 429
assert isinstance(raised.original_exception, litellm.APIError)
assert raised.original_exception.status_code == 429
assert mock_fallback.await_args.kwargs["kwargs"]["input"] == "original question"
@pytest.mark.asyncio
async def test_aresponses_fallback_uses_continuation_input_after_partial_content():
"""When output text was already streamed before the error, the fallback re-entry
must carry a continuation input with the partial assistant text instead of
retrying the original input from scratch (which would duplicate streamed content)."""
import json
from unittest.mock import Mock
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import ResponsesAPIStreamingIterator
from litellm.types.llms.openai import ErrorEvent, ErrorEventError
router = _make_router()
events = [
{"type": "response.output_text.delta", "delta": "partial answer"},
{"type": "error", "error": {"type": "server_error", "code": "internal_error", "message": "boom"}},
]
sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events)
async def mock_aiter_bytes():
yield sse_payload
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
def transform(model, parsed_chunk, logging_obj):
if parsed_chunk.get("type") == "error":
return ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(**parsed_chunk["error"]),
)
delta_event = Mock()
delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
delta_event.delta = parsed_chunk["delta"]
return delta_event
mock_config.transform_streaming_response.side_effect = transform
source = ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
fallback_event = _make_completed_event(1, 1, 2)
class _FallbackStream:
def __init__(self) -> None:
self._done = False
def __aiter__(self):
return self
async def __anext__(self):
if self._done:
raise StopAsyncIteration
self._done = True
return fallback_event
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(return_value=_FallbackStream()),
) as mock_fallback:
wrapped = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={"model": "primary", "input": "original question"},
)
collected = [ev async for ev in wrapped]
assert collected[-1] == fallback_event
raised = mock_fallback.await_args.kwargs["e"]
assert isinstance(raised, MidStreamFallbackError)
assert raised.is_pre_first_chunk is False
assert raised.generated_content == "partial answer"
continuation = mock_fallback.await_args.kwargs["kwargs"]["input"]
assert isinstance(continuation, list)
assert continuation[0]["content"][0]["text"] == "original question"
assert continuation[-2]["role"] == "developer"
assert continuation[-1]["role"] == "assistant"
assert continuation[-1]["content"][0]["text"] == "partial answer"
@pytest.mark.asyncio
async def test_aresponses_client_error_event_skips_fallback():
"""A 400-mapped in-stream error (raised as APIError, not MidStreamFallbackError)
must surface to the caller without invoking the router's fallback path."""
import litellm
router = _make_router()
class _ClientErrorSource:
completed_response = None
def __aiter__(self):
return self
async def __anext__(self):
raise litellm.APIError(
status_code=400,
message="bad request",
llm_provider="openai",
model="gpt-5",
)
wrapped = await router._aresponses_streaming_iterator(
response=_ClientErrorSource(),
initial_kwargs={"model": "primary"},
)
with patch.object(
router,
"async_function_with_fallbacks_common_utils",
new=AsyncMock(),
) as mock_fallback:
with pytest.raises(litellm.APIError) as exc_info:
async for _ in wrapped:
pass
assert exc_info.value.status_code == 400
mock_fallback.assert_not_awaited()

View file

@ -0,0 +1,357 @@
"""
Regression: in-stream error events (type="error", type="response.failed") must
raise instead of being returned as benign chunks, mirroring chat streaming
semantics (_handle_stream_fallback_error): non-retriable 4xx (except 429)
raise litellm.APIError directly; 429 and 5xx are wrapped in
MidStreamFallbackError so the Router's mid-stream fallback machinery fires.
Status mapping must consider both the OpenAI error `type` (e.g.
"invalid_request_error") and `code` (e.g. "invalid_prompt",
"rate_limit_exceeded") fields previously only `code` was read, so
type-classified client errors fell through to 500.
Also covers: ErrorEventError.param must accept dict payloads without raising a
Pydantic ValidationError (previously typed as Optional[str]).
"""
import json
import os
import sys
from unittest.mock import Mock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.responses.streaming_iterator import (
BaseResponsesAPIStreamingIterator,
ResponsesAPIStreamingIterator,
SyncResponsesAPIStreamingIterator,
)
from litellm.types.llms.openai import (
ErrorEvent,
ErrorEventError,
ResponseAPIUsage,
ResponsesAPIStreamEvents,
)
def _make_iterator() -> BaseResponsesAPIStreamingIterator:
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_config = Mock(spec=BaseResponsesAPIConfig)
mock_response = Mock()
mock_response.headers = {}
return BaseResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
def _make_error_chunk(error_type: str, code: str, message: str = "err") -> ErrorEvent:
error_obj = ErrorEventError(type=error_type, code=code, message=message)
return ErrorEvent(type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj)
def test_maybe_raise_for_error_event_wraps_unknown_error_in_mid_stream_fallback():
iterator = _make_iterator()
chunk = _make_error_chunk("server_error", "internal_error", "something went wrong")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 500
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 500
def test_maybe_raise_for_error_event_maps_rate_limit_code_to_429_mid_stream_fallback():
"""429 is retriable: it must be wrapped so the Router can fall back, carrying the mapped APIError."""
iterator = _make_iterator()
chunk = _make_error_chunk("tokens", "rate_limit_exceeded", "Too many requests")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
assert exc_info.value.generated_content == ""
assert exc_info.value.is_pre_first_chunk is True
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 429
def test_maybe_raise_for_error_event_maps_invalid_request_type_to_400():
"""Client errors classified via the `type` field must raise APIError directly (no fallback)."""
iterator = _make_iterator()
chunk = _make_error_chunk("invalid_request_error", "invalid_prompt", "bad request")
with pytest.raises(litellm.APIError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 400
assert not isinstance(exc_info.value, MidStreamFallbackError)
def test_maybe_raise_for_error_event_maps_context_length_code_to_400():
"""Client errors classified via the `code` field alone must still map to 400."""
iterator = _make_iterator()
chunk = Mock()
chunk.type = "error"
chunk.error = {"code": "context_length_exceeded", "message": "too long"}
with pytest.raises(litellm.APIError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 400
assert not isinstance(exc_info.value, MidStreamFallbackError)
def test_maybe_raise_for_error_event_maps_insufficient_quota_to_429():
"""OpenAI returns HTTP 429 for insufficient_quota; it must not map to 400 even though its type
is invalid_request_error-adjacent, and it must be wrapped for fallback."""
iterator = _make_iterator()
chunk = _make_error_chunk("invalid_request_error", "insufficient_quota", "You exceeded your current quota")
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
def test_maybe_raise_for_error_event_passes_through_normal_chunk():
iterator = _make_iterator()
chunk = Mock()
chunk.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
iterator._maybe_raise_for_error_event(chunk) # must not raise
def test_error_event_error_param_accepts_dict():
error_obj = ErrorEventError(
type="invalid_request_error",
code="context_length_exceeded",
message="too long",
param={"field": "messages", "index": 0},
)
assert isinstance(error_obj.param, dict)
def _make_async_iterator_with_events(events: list) -> ResponsesAPIStreamingIterator:
sse_payload = b"".join(f"data: {json.dumps(event)}\n\n".encode() for event in events)
async def mock_aiter_bytes():
yield sse_payload
mock_response = Mock()
mock_response.headers = {}
mock_response.aiter_bytes = mock_aiter_bytes
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
def transform(model, parsed_chunk, logging_obj):
if parsed_chunk.get("type") == "error":
return ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR,
sequence_number=0,
error=ErrorEventError(**parsed_chunk["error"]),
)
delta_event = Mock()
delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
delta_event.delta = parsed_chunk.get("delta", "")
return delta_event
mock_config.transform_streaming_response.side_effect = transform
return ResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
@pytest.mark.asyncio
async def test_async_iterator_raises_mid_stream_fallback_on_rate_limit_error_event():
iterator = _make_async_iterator_with_events(
[
{
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "rate limited"},
}
]
)
with pytest.raises(MidStreamFallbackError) as exc_info:
async for _ in iterator:
pass
assert exc_info.value.status_code == 429
assert exc_info.value.is_pre_first_chunk is True
assert exc_info.value.generated_content == ""
assert isinstance(exc_info.value.original_exception, litellm.APIError)
assert exc_info.value.original_exception.status_code == 429
@pytest.mark.asyncio
async def test_async_iterator_error_after_first_chunk_carries_generated_content():
"""An error after streamed output must expose the accumulated text so the router's
fallback can build a continuation input instead of restarting from scratch."""
iterator = _make_async_iterator_with_events(
[
{"type": "response.output_text.delta", "delta": "hello "},
{"type": "response.output_text.delta", "delta": "world"},
{
"type": "error",
"error": {"type": "server_error", "code": "internal_error", "message": "boom"},
},
]
)
chunks = []
with pytest.raises(MidStreamFallbackError) as exc_info:
async for chunk in iterator:
chunks.append(chunk)
assert len(chunks) == 2
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"
def test_maybe_raise_for_response_failed_event_with_dict_error():
"""response.failed chunks carry a dict error on .response.error; covers dict branch."""
iterator = _make_iterator()
mock_response_obj = Mock()
mock_response_obj.error = {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"}
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
def test_maybe_raise_for_error_event_null_error_obj():
"""error chunk with no error field: message and code default; wrapped as 500."""
iterator = _make_iterator()
chunk = Mock()
chunk.type = "error"
chunk.error = None
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 500
assert "Response API in-stream error" in str(exc_info.value)
def _make_failed_chunk(error: dict, usage: ResponseAPIUsage | None = None) -> Mock:
mock_response_obj = Mock()
mock_response_obj.error = error
mock_response_obj.usage = usage
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
return chunk
def test_handle_logging_failed_response_maps_rate_limit_to_429():
"""The exception logged to failure handlers must carry the mapped status, not a hardcoded 500."""
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async,
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
logged_exception = mock_run_async.call_args.kwargs["exception"]
assert isinstance(logged_exception, litellm.APIError)
assert logged_exception.status_code == 429
assert "throttled" in str(logged_exception)
def test_handle_logging_failed_response_maps_type_field_to_400():
"""Status derivation for failed-response logging must also read the error `type` field."""
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "invalid_request_error", "code": "invalid_prompt", "message": "bad prompt"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function") as mock_run_async,
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
logged_exception = mock_run_async.call_args.kwargs["exception"]
assert isinstance(logged_exception, litellm.APIError)
assert logged_exception.status_code == 400
def test_handle_logging_failed_response_records_usage_and_cost():
"""Usage on a response.failed event must reach failure spend accounting via combined_usage_object."""
iterator = _make_iterator()
usage = ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15)
chunk = _make_failed_chunk(
{"type": "server_error", "code": "server_error", "message": "boom"},
usage=usage,
)
iterator.completed_response = chunk
iterator.logging_obj._response_cost_calculator.return_value = 0.0042
with (
patch("litellm.responses.streaming_iterator.run_async_function"),
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
combined_usage = iterator.logging_obj.model_call_details["combined_usage_object"]
assert isinstance(combined_usage, litellm.Usage)
assert combined_usage.prompt_tokens == 10
assert combined_usage.completion_tokens == 5
assert combined_usage.total_tokens == 15
assert iterator.logging_obj.model_call_details["response_cost"] == 0.0042
iterator.logging_obj._response_cost_calculator.assert_called_once_with(result=chunk.response)
def test_handle_logging_failed_response_without_usage_skips_recording():
iterator = _make_iterator()
iterator.completed_response = _make_failed_chunk(
{"type": "server_error", "code": "server_error", "message": "boom"}
)
with (
patch("litellm.responses.streaming_iterator.run_async_function"),
patch("litellm.responses.streaming_iterator.executor"),
):
iterator._handle_logging_failed_response()
assert "combined_usage_object" not in iterator.logging_obj.model_call_details
iterator.logging_obj._response_cost_calculator.assert_not_called()
def test_sync_iterator_raises_mid_stream_fallback_on_rate_limit_error_event():
"""SyncResponsesAPIStreamingIterator must wrap retriable error events for fallback."""
error_payload = {
"type": "error",
"error": {"type": "tokens", "code": "rate_limit_exceeded", "message": "throttled"},
}
sse_bytes = f"data: {json.dumps(error_payload)}\n\n".encode()
mock_response = Mock()
mock_response.headers = {}
mock_response.iter_bytes.return_value = iter([sse_bytes])
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.model_call_details = {"litellm_params": {}}
mock_logging_obj.completion_start_time = None
mock_config = Mock(spec=BaseResponsesAPIConfig)
error_obj = ErrorEventError(type="tokens", code="rate_limit_exceeded", message="throttled")
mock_config.transform_streaming_response.return_value = ErrorEvent(
type=ResponsesAPIStreamEvents.ERROR, sequence_number=0, error=error_obj
)
iterator = SyncResponsesAPIStreamingIterator(
response=mock_response,
model="gpt-5",
responses_api_provider_config=mock_config,
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
)
with pytest.raises(MidStreamFallbackError) as exc_info:
for _ in iterator:
pass
assert exc_info.value.status_code == 429
assert isinstance(exc_info.value.original_exception, litellm.APIError)