Skip to content
9 changes: 6 additions & 3 deletions litellm/exceptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -464,6 +464,7 @@ def __init__(
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}"
Expand Down Expand Up @@ -507,7 +508,7 @@ def __init__(
),
)
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"
Expand Down Expand Up @@ -765,6 +766,7 @@ def __init__(
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}"
Expand All @@ -783,7 +785,7 @@ def __init__(
),
)
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):
Expand Down Expand Up @@ -815,6 +817,7 @@ def __init__(
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}"
Expand All @@ -825,7 +828,7 @@ def __init__(
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
Expand Down
5 changes: 5 additions & 0 deletions litellm/litellm_core_utils/exception_mapping_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand All @@ -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"
Expand Down Expand Up @@ -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(
Expand All @@ -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(
Expand Down
165 changes: 165 additions & 0 deletions litellm/proxy/common_utils/responses_stream_errors.py
Original file line number Diff line number Diff line change
@@ -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:
Comment thread
github-advanced-security[bot] marked this conversation as resolved.
Fixed
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"
Comment thread
cursor[bot] marked this conversation as resolved.
28 changes: 25 additions & 3 deletions litellm/proxy/proxy_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -420,6 +420,7 @@ def generate_feedback_box():
)
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,
Expand Down Expand Up @@ -8864,6 +8865,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


Expand Down Expand Up @@ -9088,10 +9090,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)
Expand Down Expand Up @@ -9202,6 +9207,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)
Expand Down Expand Up @@ -9236,8 +9243,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:
Expand All @@ -9261,8 +9273,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
Expand All @@ -9287,6 +9298,14 @@ async def async_data_generator(
e,
)
Comment thread
zoroyihan7 marked this conversation as resolved.

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):
Expand Down Expand Up @@ -9323,12 +9342,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,
)


Expand Down
6 changes: 4 additions & 2 deletions litellm/proxy/response_api_endpoints/endpoints.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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,
Expand Down
Loading
Loading