mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(responses): emit typed streaming failure events
This commit is contained in:
parent
1af7a403c6
commit
9d4bab3b70
7 changed files with 328 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
119
litellm/proxy/common_utils/responses_stream_errors.py
Normal file
119
litellm/proxy/common_utils/responses_stream_errors.py
Normal 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"
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue