mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Merge pull request #40243 from zoroyihan7/fix-responses-stream-error-events
fix(responses): emit typed streaming failure events
This commit is contained in:
commit
cda022ca68
13 changed files with 572 additions and 19 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
165
litellm/proxy/common_utils/responses_stream_errors.py
Normal file
165
litellm/proxy/common_utils/responses_stream_errors.py
Normal 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"
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue