mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
parent
5e23a5ab05
commit
249a999506
6 changed files with 691 additions and 13 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue