This commit is contained in:
zoroyihan7 2026-09-12 09:47:52 +02:00 • committed by GitHub
commit be51126847
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 413 additions and 8 deletions

View file

@ -814,6 +814,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}"
@ -824,7 +825,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

@ -0,0 +1,143 @@
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("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 _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
match details.status_code:
case HTTPStatus.UNAUTHORIZED:
return "authentication_error"
case HTTPStatus.FORBIDDEN:
return "permission_error"
case HTTPStatus.NOT_FOUND:
return "not_found_error"
case HTTPStatus.REQUEST_TIMEOUT:
return "request_timeout"
case HTTPStatus.TOO_MANY_REQUESTS:
return "rate_limit_exceeded"
case int(status) if HTTPStatus.BAD_REQUEST <= status < HTTPStatus.INTERNAL_SERVER_ERROR:
return "invalid_request_error"
case _:
return "server_error"
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 = _FailureDetails.model_validate(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

@ -400,6 +400,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,
@ -8957,10 +8958,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)
@ -9071,6 +9075,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)
@ -9105,8 +9111,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:
@ -9156,6 +9167,12 @@ 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
return
if isinstance(e, HTTPException):
raise e
elif isinstance(e, StreamingCallbackError):
@ -9192,12 +9209,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

@ -569,6 +569,11 @@ class BaseResponsesAPIStreamingIterator:
message=error_message,
llm_provider=self.custom_llm_provider or "",
model=self.model or "",
body={ # mutable-ok: OpenAI APIError reads code/type only from a dict body
"code": error_code,
"type": error_type,
"message": error_message,
},
)
if 400 <= status_code < 500 and status_code != 429:
raise mapped_exception

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

@ -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,127 @@ 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)
@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", 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 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, (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,107 @@ 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"),
("/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"],
) -> 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",
}[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 == "response_failed" 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")
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].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_exceeded" if error_kind in ("rate_limit", "numeric_rate_limit") else "server_error"
)
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")