mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(responses): preserve failure metadata at streaming boundaries
This commit is contained in:
parent
9d4bab3b70
commit
089703d10b
4 changed files with 114 additions and 31 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue