fix(responses): preserve failure metadata at streaming boundaries

This commit is contained in:
zoroyihan7 2026-09-08 11:02:38 +00:00
parent 9d4bab3b70
commit 089703d10b
4 changed files with 114 additions and 31 deletions

View file

@ -1,9 +1,10 @@
import time
from collections.abc import Mapping
from http import HTTPStatus
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict
from pydantic import BaseModel, ConfigDict, field_validator
from litellm._logging import redact_internal_details_from_client_message
from litellm._uuid import uuid
@ -35,6 +36,16 @@ class _FailureDetails(BaseModel):
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:
if isinstance(exception, MidStreamFallbackError) and exception.original_exception is not None:
@ -50,9 +61,21 @@ def _response_error_code(details: _FailureDetails) -> str:
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"
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:
@ -62,16 +85,15 @@ class ResponsesStreamErrorState:
self.created_at: int | None = None
self.sequence_number = -1
self.terminal_emitted = False
self._pending_event: _StreamEvent | None = None
@staticmethod
def observe_chunk(chunk: object) -> _StreamEvent | None:
if not isinstance(chunk, (BaseModel, Mapping)):
return None
return _StreamEvent.model_validate(chunk)
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, event: _StreamEvent | None) -> None:
def mark_emitted(self, frame: str | bytes) -> str | bytes:
event: Final = self._pending_event
if event is None:
return
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:
@ -81,6 +103,7 @@ class ResponsesStreamErrorState:
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:

View file

@ -8841,7 +8841,8 @@ 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
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)
@ -8876,10 +8877,10 @@ async def async_data_generator(
if not raw_passthrough:
try:
formatted_chunk: Final = _format_streaming_sse_chunk(chunk=chunk)
if error_state is not None:
error_state.mark_emitted(responses_event)
yield formatted_chunk
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

View file

@ -21,9 +21,11 @@ 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
@ -882,9 +884,44 @@ async def test_async_data_generator_mid_stream_exception_yields_error_payload(
@pytest.mark.asyncio
@pytest.mark.parametrize("terminal", ["completed", "serialization_failure", "failure_after_completed"])
@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"],
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"]
@ -910,16 +947,23 @@ async def test_responses_stream_keeps_tool_deltas_and_only_emits_a_valid_termina
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")
raise litellm.APIError(
status_code=500, message="Stream close failed", llm_provider="openai", model="gpt-6-astra"
)
frames: Final = [
frame
@ -940,14 +984,19 @@ async def test_responses_stream_keeps_tool_deltas_and_only_emits_a_valid_termina
assert payloads[0]["response"]["id"] == "resp_visible"
assert payloads[1] == tool_delta.model_dump()
assert len(payloads) == 3
if terminal == "serialization_failure":
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"] == "server_error"
assert "serialize" in failure.response.error["message"].lower()
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"

View file

@ -61,7 +61,9 @@ async def test_streaming_upstream_errors_keep_the_client_protocol(
"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)
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(
@ -78,8 +80,10 @@ async def test_streaming_upstream_errors_keep_the_client_protocol(
)
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"})},
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(
@ -91,12 +95,18 @@ async def test_streaming_upstream_errors_keep_the_client_protocol(
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"]
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"