Merge pull request #40243 from zoroyihan7/fix-responses-stream-error-events

fix(responses): emit typed streaming failure events
This commit is contained in:
Mateo Wang 2026-09-18 17:29:57 -07:00 • committed by GitHub
commit cda022ca68
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 572 additions and 19 deletions

View file

@ -464,6 +464,7 @@ class RateLimitError(openai.RateLimitError):
rate_limit_type: str | RateLimitType | None = None,
headers: dict[str, str] | None = None,
detail: Any = None,
body: object | None = None,
):
self.status_code = 429
self.message = f"litellm.RateLimitError: {message}"
@ -507,7 +508,7 @@ class RateLimitError(openai.RateLimitError):
),
)
super().__init__(
self.message, response=self.response, body=None
self.message, response=self.response, body=body
) # Call the base class constructor with the parameters it needs
self.code = "429"
self.type = "throttling_error"
@ -765,6 +766,7 @@ class InternalServerError(openai.InternalServerError):
litellm_debug_info: str | None = None,
max_retries: int | None = None,
num_retries: int | None = None,
body: object | None = None,
):
self.status_code = 500
self.message = f"litellm.InternalServerError: {message}"
@ -783,7 +785,7 @@ class InternalServerError(openai.InternalServerError):
),
)
super().__init__(
self.message, response=self.response, body=None
self.message, response=self.response, body=body
) # Call the base class constructor with the parameters it needs
def __str__(self):
@ -815,6 +817,7 @@ class APIError(openai.APIError):
litellm_debug_info: str | None = None,
max_retries: int | None = None,
num_retries: int | None = None,
body: object | None = None,
):
self.status_code = status_code
self.message = f"litellm.APIError: {message}"
@ -825,7 +828,7 @@ class APIError(openai.APIError):
self.num_retries = num_retries
if request is None:
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
super().__init__(self.message, request=request, body=None)
super().__init__(self.message, request=request, body=body)
def __str__(self):
_message = self.message

View file

@ -307,6 +307,7 @@ def _map_openai_exception(
model=model,
llm_provider=custom_llm_provider,
response=response,
body=getattr(original_exception, "body", None),
)
elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str):
raise ContextWindowExceededError(
@ -381,6 +382,7 @@ def _map_openai_exception(
message=f"{exception_provider} - {message}",
model=model,
llm_provider=custom_llm_provider,
body=getattr(original_exception, "body", None),
)
elif "Request too large" in error_str:
raise RateLimitError(
@ -389,6 +391,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif (
"The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable"
@ -460,6 +463,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif original_exception.status_code == 500:
raise InternalServerError(
@ -468,6 +472,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif original_exception.status_code == 502:
raise BadGatewayError(

View file

@ -0,0 +1,165 @@
import time
from collections.abc import Mapping
from http import HTTPStatus
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict, field_validator
from litellm._logging import redact_internal_details_from_client_message
from litellm._uuid import uuid
from litellm.exceptions import MidStreamFallbackError
from litellm.types.llms.openai import ResponseFailedEvent, ResponsesAPIResponse, ResponsesAPIStreamEvents
class _ResponseIdentity(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
id: str | None = None
model: str | None = None
created_at: int | None = None
class _StreamEvent(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
type: str | None = None
sequence_number: int | None = None
response: _ResponseIdentity | None = None
class _FailureDetails(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
message: str | None = None
code: str | int | None = None
type: str | None = None
status_code: int | None = None
@field_validator("message", mode="before")
@classmethod
def normalize_message(cls, value: object) -> str | None:
return value if isinstance(value, str) else None
@field_validator("code", mode="before")
@classmethod
def normalize_code(cls, value: object) -> str | int | None:
return value if isinstance(value, (str, int)) and not isinstance(value, bool) else None
@field_validator("type", mode="before")
@classmethod
def normalize_type(cls, value: object) -> str | None:
return value if isinstance(value, str) else None
def _original_failure(exception: Exception) -> Exception:
current = exception # rebind-ok: the recursion gate requires iterative wrapper traversal
while isinstance(current, MidStreamFallbackError) and current.original_exception is not None:
current = current.original_exception
return current
def _failure_details(original: Exception) -> _FailureDetails:
mapped: Final = _FailureDetails.model_validate(original)
body: Final = getattr(original, "body", None)
if not isinstance(body, Mapping):
return mapped
upstream: Final = _FailureDetails.model_validate(body)
return _FailureDetails(
message=upstream.message or mapped.message,
code=upstream.code if upstream.code is not None else mapped.code,
type=upstream.type or mapped.type,
status_code=mapped.status_code,
)
_CLIENT_ERROR_CODES: Final = MappingProxyType(
{
int(HTTPStatus.UNAUTHORIZED): "authentication_error",
int(HTTPStatus.FORBIDDEN): "permission_error",
int(HTTPStatus.NOT_FOUND): "not_found_error",
int(HTTPStatus.REQUEST_TIMEOUT): "request_timeout",
int(HTTPStatus.TOO_MANY_REQUESTS): "rate_limit_exceeded",
}
)
def _status_error_code(status_code: int | None) -> str:
if status_code is None or not HTTPStatus.BAD_REQUEST <= status_code < HTTPStatus.INTERNAL_SERVER_ERROR:
return "server_error"
return _CLIENT_ERROR_CODES.get(status_code, "invalid_request_error")
def _response_error_code(details: _FailureDetails) -> str:
for value in (details.code, details.type):
if value == "insufficient_quota":
return "insufficient_quota"
if value in (429, "429") or isinstance(value, str) and value.startswith("rate_limit"):
return "rate_limit_exceeded"
if isinstance(details.code, str) and details.code and not details.code.isdecimal():
return details.code
return _status_error_code(details.status_code)
class ResponsesStreamErrorState:
def __init__(self) -> None:
self.response_id: str | None = None
self.model: str | None = None
self.created_at: int | None = None
self.sequence_number = -1
self.terminal_emitted = False
self._pending_event: _StreamEvent | None = None
def observe_chunk(self, chunk: object) -> None:
self._pending_event = _StreamEvent.model_validate(chunk) if isinstance(chunk, (BaseModel, Mapping)) else None
def mark_emitted(self, frame: str | bytes) -> str | bytes:
event: Final = self._pending_event
if event is None:
return frame
if event.sequence_number is not None:
self.sequence_number = max(self.sequence_number, event.sequence_number)
if event.response is not None:
self.response_id = event.response.id or self.response_id
self.model = event.response.model or self.model
if event.response.created_at is not None:
self.created_at = event.response.created_at
if event.type in ("response.completed", "response.failed", "response.incomplete"):
self.terminal_emitted = True
return frame
def format_failure(self, exception: Exception) -> str | None:
if self.terminal_emitted:
return None
original: Final = _original_failure(exception)
details: Final = _failure_details(original)
response: Final = ResponsesAPIResponse.model_validate(
MappingProxyType(
{
"id": self.response_id or f"resp_{uuid.uuid4().hex}",
"object": "response",
"created_at": self.created_at if self.created_at is not None else int(time.time()),
"model": self.model,
"status": "failed",
"output": (),
"error": MappingProxyType(
{
"code": _response_error_code(details),
"message": redact_internal_details_from_client_message(details.message or str(original)),
}
),
}
)
)
event: Final = ResponseFailedEvent.model_validate(
MappingProxyType(
{
"type": ResponsesAPIStreamEvents.RESPONSE_FAILED,
"response": response,
"sequence_number": self.sequence_number + 1,
}
)
)
payload: Final = event.model_dump_json(exclude_none=True)
self.terminal_emitted = True
return f"event: response.failed\ndata: {payload}\n\n"

View file

@ -421,6 +421,7 @@ from litellm.proxy.common_utils.periodic_reload_schedule import (
)
from litellm.proxy.common_utils.proxy_state import ProxyState
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.responses_stream_errors import ResponsesStreamErrorState
from litellm.proxy.common_utils.scheduled_job_stagger import (
apply_scheduled_job_stagger,
attach_job_timing_logger,
@ -8894,6 +8895,7 @@ def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes:
_SSE_FRAME_DELIMITERS: Final = ("\r\n\r\n", "\n\n", "\r\r")
_OPENAI_STREAM_DONE_FRAME: Final = "data: [DONE]\n\n"
_MAX_RAW_SSE_BUFFER_CHARS: Final = 8 * 1024 * 1024
@ -9118,10 +9120,13 @@ async def async_data_generator(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
*,
responses_stream_errors: bool = False,
):
verbose_proxy_logger.debug("inside generator")
stream_completed = False
client_disconnected = False
error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None
try:
error_message: str | None = None
requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data)
@ -9232,6 +9237,8 @@ async def async_data_generator(
fallback_metadata_event_sent = True
continue
if error_state is not None:
error_state.observe_chunk(cast(object, chunk)) # cast-ok: the helper validates legacy untyped chunks
raw_passthrough = False
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
@ -9266,8 +9273,13 @@ async def async_data_generator(
if not raw_passthrough:
try:
yield _format_streaming_sse_chunk(chunk=chunk)
if error_state is not None:
yield error_state.mark_emitted(_format_streaming_sse_chunk(chunk=chunk))
else:
yield _format_streaming_sse_chunk(chunk=chunk)
except Exception as e:
if error_state is not None:
raise
yield f"data: {e}\n\n"
if pending_fallback_event:
@ -9291,8 +9303,7 @@ async def async_data_generator(
yield error_message
# OpenAI-compatible streams terminate with data: [DONE]; Google GenAI (?alt=sse) does not.
if not request_data.get("_litellm_skip_openai_stream_done"):
done_message: Final = "[DONE]"
yield f"data: {done_message}\n\n"
yield _OPENAI_STREAM_DONE_FRAME
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit are
# BaseException, so they bypass the success/failure logging callbacks
@ -9317,6 +9328,14 @@ async def async_data_generator(
e,
)
if error_state is not None:
stream_completed = True
error_frame: Final = error_state.format_failure(e)
if error_frame is not None:
yield error_frame
if not request_data.get("_litellm_skip_openai_stream_done"):
yield _OPENAI_STREAM_DONE_FRAME
return
if isinstance(e, HTTPException):
raise e
elif isinstance(e, StreamingCallbackError):
@ -9353,12 +9372,15 @@ def select_data_generator(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
*,
responses_stream_errors: bool = False,
):
return async_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
responses_stream_errors=responses_stream_errors,
)

View file

@ -3,6 +3,7 @@ import json
import time
from collections.abc import AsyncIterator, Awaitable, Mapping
from enum import Enum
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args
from uuid import uuid4
@ -243,6 +244,7 @@ async def responses_api(
version,
)
native_data_generator: Final = partial(select_data_generator, responses_stream_errors=True)
data = await _read_request_body(request=request)
# Check if polling via cache should be used for this request
@ -329,7 +331,7 @@ async def responses_api(
llm_router=llm_router,
proxy_config=proxy_config,
proxy_logging_obj=proxy_logging_obj,
select_data_generator=select_data_generator,
select_data_generator=native_data_generator,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
@ -355,7 +357,7 @@ async def responses_api(
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
select_data_generator=native_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,

View file

@ -75,6 +75,10 @@ class _StreamEventParser:
parse: Callable[[str], _StreamEvent] = staticmethod(json.loads)
def _sse_frame_data(frame: str) -> str | None:
return next((line[6:].strip() for line in frame.splitlines() if line.startswith("data: ")), None)
async def _never_receive() -> Message:
await asyncio.Event().wait()
raise AssertionError("unreachable")
@ -224,8 +228,7 @@ async def background_streaming_task(
if isinstance(chunk, bytes):
chunk = chunk.decode("utf-8")
if isinstance(chunk, str) and chunk.startswith("data: "):
chunk_data = chunk[6:].strip()
if isinstance(chunk, str) and (chunk_data := _sse_frame_data(chunk)) is not None:
if chunk_data == "[DONE]":
break

View file

@ -212,18 +212,21 @@ def _error_event_fields(error_obj: object) -> tuple[str, str | None, str | None]
raw_code = None
message: Final = str(raw_message) if raw_message is not None else "Response API in-stream error"
error_type: Final = raw_type if isinstance(raw_type, str) else None
code: Final = raw_code if isinstance(raw_code, str) else None
code: Final = str(raw_code) if isinstance(raw_code, (str, int)) and not isinstance(raw_code, bool) else None
return message, error_type, code
def _status_code_for_error_field(field: str) -> int | None:
if field.isdecimal() and 400 <= int(field) <= 599:
return int(field)
return _ERROR_CODE_HTTP_STATUS.get(field)
def _status_code_for_error_fields(error_type: str | None, error_code: str | None) -> int:
fields: Final = tuple(field for field in (error_code, error_type) if field is not None)
if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields):
return 429
return next(
(_ERROR_CODE_HTTP_STATUS[field] for field in fields if field in _ERROR_CODE_HTTP_STATUS),
500,
)
return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500)
def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool:

View file

@ -1482,6 +1482,44 @@ class TestBackgroundStreamingTerminalEvents:
assert final_call.kwargs["status"] == "failed"
assert final_call.kwargs["error"] == error_payload
@pytest.mark.asyncio
async def test_named_event_failed_frame_sets_failed_status_and_error(self):
from litellm.proxy.response_polling.background_streaming import (
background_streaming_task,
)
error_payload = {
"code": "cyber_policy",
"message": "Your request was flagged for possible cybersecurity risk and was not completed",
}
failed_event = {
"type": "response.failed",
"sequence_number": 5,
"response": {"id": "resp_123", "status": "failed", "error": error_payload, "output": []},
}
async def _body_iterator():
yield b'data: {"type": "response.in_progress"}\n\n'
yield f"event: response.failed\ndata: {json.dumps(failed_event)}\n\n".encode()
yield b"data: [DONE]\n\n"
mock_response = Mock()
mock_response.body_iterator = _body_iterator()
handler = AsyncMock(spec=ResponsePollingHandler)
kwargs = _make_background_streaming_kwargs("poll_named_event", handler)
with patch( # test-quality-ok: the processor is built inside the task, same idiom as the sibling tests
"litellm.proxy.response_polling.background_streaming.ProxyBaseLLMRequestProcessing"
) as MockProcessor:
MockProcessor.return_value.base_process_llm_request = AsyncMock(
return_value=mock_response
)
await background_streaming_task(**kwargs)
final_call = handler.update_state.call_args_list[-1]
assert final_call.kwargs["status"] == "failed"
assert final_call.kwargs["error"] == error_payload
@pytest.mark.asyncio
async def test_response_incomplete_sets_incomplete_status_and_details(self):
"""Test that a response.incomplete stream event results in incomplete status"""

View file

@ -130,7 +130,6 @@ class TestPollingEndpointPreCallGuard:
"litellm.proxy.proxy_server.proxy_config": MagicMock(),
"litellm.proxy.proxy_server.proxy_logging_obj": AsyncMock(),
"litellm.proxy.proxy_server.redis_usage_cache": AsyncMock(),
"litellm.proxy.proxy_server.select_data_generator": None,
"litellm.proxy.proxy_server.user_api_base": None,
"litellm.proxy.proxy_server.user_max_tokens": None,
"litellm.proxy.proxy_server.user_model": None,

View file

@ -1437,6 +1437,29 @@ def test_openai_compatible_vendor_400_keeps_body_but_not_headers():
assert not exc_info.value.response.headers
@pytest.mark.parametrize(
("status_code", "mapped_class"), [(429, litellm.RateLimitError), (500, litellm.InternalServerError)]
)
def test_openai_429_and_500_keep_body(status_code: int, mapped_class: type[openai.APIError]):
with pytest.raises(mapped_class) as exc_info:
exception_type(
model="gpt-5.4-mini",
original_exception=_openai_handler_error(
"server_error", {}, status_code=status_code, message="upstream cannot complete this response"
),
custom_llm_provider="openai",
completion_kwargs={},
extra_kwargs={},
)
assert exc_info.value.body == {
**_GUARDRAIL_BLOCK_ERROR,
"type": "server_error",
"code": str(status_code),
"message": "upstream cannot complete this response",
}
def test_litellm_proxy_repeated_response_header_keeps_each_value():
repeated = [("x-litellm-call-id", "call-guardrail"), ("set-cookie", "a=1"), ("set-cookie", "b=2")]

View file

@ -17,15 +17,20 @@ from __future__ import annotations
import asyncio
import json
from collections.abc import AsyncIterator
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import Response
from fastapi import HTTPException, Response
from fastapi.responses import StreamingResponse
from openai import APIError as OpenAIAPIError
from pydantic import BaseModel
import litellm
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
import litellm.proxy.proxy_server as ps
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import (
_apply_streaming_chunk_hooks,
@ -42,6 +47,12 @@ from litellm.proxy.proxy_server import (
data_generator,
select_data_generator,
)
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponseCreatedEvent,
ResponseFailedEvent,
ResponsesAPIResponse,
)
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage
from .conftest import normalize
@ -872,6 +883,145 @@ async def test_async_data_generator_mid_stream_exception_yields_error_payload(
assert any(isinstance(item, str) and item.startswith('data: {"error":') for item in out)
_UPSTREAM_BODY: Final = {
"code": "cyber_policy",
"message": "Upstream rejected request: flagged for possible cybersecurity risk",
"type": None,
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"terminal,upstream_error,expected_code",
[
("completed", None, None),
("serialization_failure", None, "server_error"),
("failure_after_completed", None, None),
pytest.param(
"upstream_failure",
litellm.AuthenticationError(
message="Upstream rejected request", llm_provider="openai", model="gpt-6-astra"
),
"authentication_error", id="authentication_error",
),
pytest.param(
"upstream_failure",
OpenAIAPIError(
message="Upstream rejected request",
request=httpx.Request("POST", "https://streaming.example/v1/responses"),
body={"code": {"reason": "overloaded"}, "type": {"unexpected": "object"}},
),
"server_error", id="structured_provider_error_fields",
),
pytest.param(
"upstream_failure",
litellm.InternalServerError(
message="Upstream rejected request", llm_provider="openai", model="gpt-6-astra", body=_UPSTREAM_BODY
),
"cyber_policy", id="upstream_body_code_and_message",
),
*(
pytest.param(
"upstream_failure", HTTPException(status_code=status, detail="Upstream rejected request"),
code, id=f"http_{status}",
)
for status, code in (
(400, "invalid_request_error"), (403, "permission_error"), (404, "not_found_error"),
(408, "request_timeout"), (422, "invalid_request_error"), (500, "server_error"), (503, "server_error"),
)
),
],
)
async def test_responses_stream_keeps_tool_deltas_and_only_emits_a_valid_terminal(
terminal: Literal["completed", "serialization_failure", "failure_after_completed", "upstream_failure"],
upstream_error: HTTPException | OpenAIAPIError | None,
expected_code: str | None,
) -> None:
class ToolDelta(BaseModel):
type: Literal["response.function_call_arguments.delta"]
sequence_number: int
item_id: str
output_index: int
delta: str
class UnserializableTerminal(BaseModel):
type: Literal["response.completed"]
sequence_number: int
response: ResponsesAPIResponse
invalid: object
response: Final = ResponsesAPIResponse(id="resp_visible", created_at=1, model="gpt-6-astra", output=[])
created: Final = ResponseCreatedEvent.model_validate(
{"type": "response.created", "sequence_number": 0, "response": response}
)
completed: Final = ResponseCompletedEvent.model_validate(
{"type": "response.completed", "sequence_number": 2, "response": response}
)
tool_delta: Final = ToolDelta(
type="response.function_call_arguments.delta", sequence_number=1, item_id="fc_stream_error",
output_index=0, delta='{"path":"partial',
)
original_status: Final = (
upstream_error.status_code if isinstance(upstream_error, (HTTPException, litellm.AuthenticationError)) else None
)
async def upstream() -> AsyncIterator[BaseModel]:
yield created
yield tool_delta
if upstream_error is not None:
raise upstream_error
yield (
UnserializableTerminal(type="response.completed", sequence_number=2, response=response, invalid=object())
if terminal == "serialization_failure" else completed
)
if terminal == "failure_after_completed":
raise litellm.APIError(
status_code=500, message="Stream close failed", llm_provider="openai", model="gpt-6-astra"
)
frames: Final = [
frame
async for frame in select_data_generator(
response=upstream(),
user_api_key_dict=_user_auth(),
request_data={},
responses_stream_errors=True,
)
]
decoded: Final = tuple(frame.decode() if isinstance(frame, bytes) else frame for frame in frames)
event_frames: Final = tuple(frame for frame in decoded if frame != "data: [DONE]\n\n")
payloads: Final = tuple(
json.loads(next(line[6:] for line in frame.splitlines() if line.startswith("data: ")))
for frame in event_frames
)
assert decoded[-1] == "data: [DONE]\n\n"
assert len(decoded) == len(event_frames) + 1
assert payloads[0]["response"]["id"] == "resp_visible"
assert payloads[1] == tool_delta.model_dump()
assert len(payloads) == 3
if terminal in ("serialization_failure", "upstream_failure"):
failure: Final = ResponseFailedEvent.model_validate(payloads[-1])
assert event_frames[-1].startswith("event: response.failed\n")
assert failure.response.id == "resp_visible"
assert failure.response.status == "failed"
assert failure.response.error is not None
assert failure.response.error["code"] == expected_code
if upstream_error is None:
assert "serialize" in failure.response.error["message"].lower()
else:
assert "Upstream rejected request" in failure.response.error["message"]
if isinstance(upstream_error, litellm.InternalServerError):
assert failure.response.error["message"] == _UPSTREAM_BODY["message"]
if isinstance(upstream_error, (HTTPException, litellm.AuthenticationError)):
assert upstream_error.status_code == original_status
assert payloads[-1]["sequence_number"] > payloads[1]["sequence_number"]
else:
assert payloads[-1]["type"] == "response.completed"
assert payloads[-1]["sequence_number"] == 2
assert "error" not in payloads[-1]
# ---------------------------------------------------------------------------
# select_data_generator
# ---------------------------------------------------------------------------

View file

@ -3,10 +3,12 @@ Test for response_api_endpoints/endpoints.py
"""
import unittest
from typing import Any
from typing import Any, Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from fastapi.testclient import TestClient
from httpx import Response
@ -14,6 +16,111 @@ import litellm
from litellm.proxy.proxy_server import app
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path,error_kind",
[
("/v1/responses", "rate_limit"),
("/v1/responses", "numeric_rate_limit"),
("/v1/responses", "server_error"),
("/v1/responses", "response_failed"),
("/v1/responses", "cyber_policy"),
("/cursor/chat/completions", "server_error"),
("/v1/chat/completions", "server_error"),
],
)
async def test_streaming_upstream_errors_keep_the_client_protocol(
monkeypatch: pytest.MonkeyPatch,
path: str,
error_kind: Literal["rate_limit", "numeric_rate_limit", "server_error", "response_failed", "cyber_policy"],
) -> None:
import litellm.proxy.proxy_server as ps
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
model: Final = "gpt-6-astra"
message: Final = "Upstream cannot complete this response"
code: Final = {
"rate_limit": "rate_limit_exceeded", "numeric_rate_limit": "429",
"server_error": "server_error", "response_failed": "server_error", "cyber_policy": "cyber_policy",
}[error_kind]
error: Final = {"message": message, "code": code, "type": None, "param": "input"}
response: Final = {"id": "resp_upstream", "object": "response", "created_at": 1,
"status": "in_progress", "model": model, "output": [],
"parallel_tool_calls": True, "tool_choice": "auto", "tools": []}
created: Final = {"type": "response.created", "sequence_number": 0, "response": response}
tool_added: Final = {"type": "response.output_item.added", "sequence_number": 1, "output_index": 0,
"item": {"type": "function_call", "id": "fc_partial", "call_id": "call_partial",
"name": "read_file", "arguments": "", "status": "in_progress"}}
tool_delta: Final = {"type": "response.function_call_arguments.delta", "sequence_number": 2,
"item_id": "fc_partial", "output_index": 0, "delta": '{"path":"partial'}
failed: Final = (
{"type": "response.failed", "sequence_number": 9,
"response": {**response, "status": "failed", "error": error}}
if error_kind in ("response_failed", "cyber_policy") else {"type": "error", "error": error}
)
chat: Final = {"id": "chatcmpl_partial", "object": "chat.completion.chunk", "created": 1,
"model": model, "choices": [{"index": 0, "delta": {"content": "partial"},
"finish_reason": None}]}
is_chat: Final = path == "/v1/chat/completions"
partial: Final = path != "/v1/responses" or error_kind in ("numeric_rate_limit", "response_failed", "cyber_policy")
response_events: Final = (created, tool_added, tool_delta, failed) if partial else (failed,)
upstream_events: Final = (chat, {"error": error}) if is_chat else response_events
wire: Final = "".join("data: " + json.dumps(event) + "\n\n" for event in upstream_events)
upstream_url: Final = "https://streaming.example/v1"
router: Final = litellm.Router(
model_list=[{"model_name": model, "litellm_params": {
"model": "openai/" + model, "api_base": upstream_url, "api_key": "fixture-key"}}],
num_retries=0,
)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, _auth_override)
with respx.mock as transport:
transport.post(upstream_url + ("/chat/completions" if is_chat else "/responses")).respond(
200, content=wire, headers={"Content-Type": "text/event-stream"}
)
async with httpx.AsyncClient(transport=httpx.ASGITransport(app), base_url="http://testserver") as client:
result: Final = await client.post(
path, json={
"model": model, "stream": True,
**({"messages": [{"role": "user", "content": "hello"}]} if is_chat else {"input": "hello"}),
},
)
frames: Final = tuple(frame for frame in result.text.split("\n\n") if "data: " in frame)
events: Final = tuple(
json.loads(next(line[6:] for line in frame.splitlines() if line.startswith("data: ")))
for frame in frames if "data: [DONE]" not in frame
)
assert result.status_code == 200, result.text
assert message in result.text
if path == "/v1/responses":
assert frames[-1] == "data: [DONE]", result.text
assert frames[-2].startswith("event: response.failed\n"), result.text
if partial:
assert [event["type"] for event in events] == [
"response.created", "response.output_item.added",
"response.function_call_arguments.delta", "response.failed",
]
assert events[2]["delta"] == tool_delta["delta"]
assert events[-1]["sequence_number"] == events[-2]["sequence_number"] + 1
assert events[-1]["response"]["id"] == events[0]["response"]["id"]
else:
assert [event["type"] for event in events] == ["response.failed"]
assert events[0]["sequence_number"] == 0
assert events[0]["response"]["id"].startswith("resp_")
assert events[-1]["response"]["status"] == "failed"
assert events[-1]["response"]["error"]["code"] == {
"rate_limit": "rate_limit_exceeded", "numeric_rate_limit": "rate_limit_exceeded",
"server_error": "server_error", "response_failed": "server_error", "cyber_policy": "cyber_policy",
}[error_kind]
assert events[-1]["response"]["error"]["message"] == message
else:
assert events[0]["object"] == "chat.completion.chunk", result.text
assert "response.failed" not in result.text
assert "error" in events[-1]
class TestResponsesAPIEndpoints(unittest.TestCase):
@pytest.mark.asyncio
@patch("litellm.proxy.proxy_server.llm_router")

View file

@ -340,6 +340,36 @@ def test_maybe_raise_for_response_failed_event_with_dict_error():
assert exc_info.value.status_code == 429
@pytest.mark.parametrize("code", [429, "429"])
def test_response_failed_numeric_code_maps_to_its_http_status(code: int | str):
iterator = _make_iterator()
mock_response_obj = Mock()
mock_response_obj.error = {"code": code, "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
assert isinstance(exc_info.value.original_exception, litellm.RateLimitError)
def test_response_failed_unknown_code_keeps_upstream_code_and_message_on_mapped_exception():
iterator = _make_iterator()
upstream_message = "This content was flagged for possible cybersecurity risk."
mock_response_obj = Mock()
mock_response_obj.error = {"code": "cyber_policy", "message": upstream_message}
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)
mapped = exc_info.value.original_exception
assert isinstance(mapped, litellm.InternalServerError)
assert mapped.code == "cyber_policy"
assert mapped.body == {"message": upstream_message, "type": None, "code": "cyber_policy"}
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()
@ -523,6 +553,9 @@ def test_every_openai_sdk_response_error_code_has_explicit_status_mapping():
("failed_to_download_image", 400),
("image_file_not_found", 400),
("totally_unknown_future_code", 500),
("429", 429),
("503", 503),
("200", 500),
],
)
def test_status_code_for_documented_response_error_codes(code: str, expected_status: int):