fix(responses): emit typed streaming failure events

This commit is contained in:
zoroyihan7 2026-09-08 10:46:16 +00:00
parent 1af7a403c6
commit 9d4bab3b70
7 changed files with 328 additions and 6 deletions

View file

@ -789,6 +789,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}"
@ -799,7 +800,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,119 @@
import time
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict
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
def _original_failure(exception: Exception) -> Exception:
if isinstance(exception, MidStreamFallbackError) and exception.original_exception is not None:
return _original_failure(exception.original_exception)
return exception
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
if details.status_code == 429:
return "rate_limit_exceeded"
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
@staticmethod
def observe_chunk(chunk: object) -> _StreamEvent | None:
if not isinstance(chunk, (BaseModel, Mapping)):
return None
return _StreamEvent.model_validate(chunk)
def mark_emitted(self, event: _StreamEvent | None) -> None:
if event is None:
return
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
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

@ -385,6 +385,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,
@ -8723,10 +8724,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)
@ -8837,6 +8841,7 @@ async def async_data_generator(
fallback_metadata_event_sent = True
continue
responses_event: Final = error_state.observe_chunk(chunk) if error_state is not None else None
raw_passthrough = False
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
@ -8871,8 +8876,13 @@ async def async_data_generator(
if not raw_passthrough:
try:
yield _format_streaming_sse_chunk(chunk=chunk)
formatted_chunk: Final = _format_streaming_sse_chunk(chunk=chunk)
if error_state is not None:
error_state.mark_emitted(responses_event)
yield formatted_chunk
except Exception as e:
if error_state is not None:
raise
yield f"data: {e}\n\n"
if pending_fallback_event:
@ -8922,6 +8932,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):
@ -8958,12 +8974,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

@ -579,6 +579,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

@ -17,15 +17,18 @@ 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 pytest
from fastapi import Response
from fastapi.responses import StreamingResponse
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 +45,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 +881,80 @@ 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", ["completed", "serialization_failure", "failure_after_completed"])
async def test_responses_stream_keeps_tool_deltas_and_only_emits_a_valid_terminal(
terminal: Literal["completed", "serialization_failure", "failure_after_completed"],
) -> 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',
)
async def upstream() -> AsyncIterator[BaseModel]:
yield created
yield tool_delta
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 == "serialization_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"] == "server_error"
assert "serialize" in failure.response.error["message"].lower()
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,97 @@ 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"
upstream_events: Final = (chat, {"error": error}) if is_chat else (created, tool_added, tool_delta, failed)
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
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"]
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")