diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 43611eb3192..a78e6f17cc9 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -2303,9 +2303,29 @@ class CustomStreamWrapper: 429 (rate-limit) is explicitly exempted from the 4xx filter because it is transient and the Router should switch to another model group. + + An error an inner stream already wrapped (the chat-to-Responses bridge + consumes a Responses stream) is rebuilt around the provider exception + with this wrapper's own bookkeeping, so the Router's one-level unwrap + surfaces the provider exception and is_pre_first_chunk says whether + this wrapper's consumer received anything (the inner stream counts a + lifecycle event the bridge never forwards as its first chunk). """ from litellm.exceptions import MidStreamFallbackError + if isinstance(e, MidStreamFallbackError): + self._restore_consumer_correlation_context() + if e.original_exception is None: + raise e + raise MidStreamFallbackError( + message=str(e.original_exception), + model=self.model, + llm_provider=self.custom_llm_provider or "anthropic", + original_exception=e.original_exception, + generated_content=self.response_uptil_now, + is_pre_first_chunk=not self.sent_first_chunk, + ) + # Map to OpenAI exception format. Some providers' mappers (e.g. # _map_anthropic_exception, _map_aleph_alpha_exception) synchronously # log a debug diagnostic (the raw status code) as part of mapping - diff --git a/tests/integration/_support/responses_stream.py b/tests/integration/_support/responses_stream.py new file mode 100644 index 00000000000..1bcfebfba90 --- /dev/null +++ b/tests/integration/_support/responses_stream.py @@ -0,0 +1,128 @@ +import json +from collections.abc import Callable, Iterator, Mapping, Sequence +from typing import Final + +from integration._support.client import object_value, string_value +from integration._support.wire import Reply, Request +from pydantic import JsonValue + +AZURE_TARGET: Final = "/openai/v1/responses?api-version=" +OPENAI_TARGET: Final = "/responses" +RATE_LIMIT_MESSAGE: Final = "Your requests to gpt-6 have exceeded token rate limit." + + +def frame(event: Mapping[str, JsonValue]) -> bytes: + return f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() + + +def response_object(identity: str, status: str, **fields: JsonValue) -> dict[str, JsonValue]: + return { + "id": identity, + "object": "response", + "created_at": 1, + "status": status, + "model": "gpt-6", + "output": [], + "usage": None, + **fields, + } + + +def created(identity: str) -> Mapping[str, JsonValue]: + return {"type": "response.created", "sequence_number": 0, "response": response_object(identity, "in_progress")} + + +def error_event(error: Mapping[str, JsonValue] | None) -> Mapping[str, JsonValue]: + return {"type": "error", "sequence_number": 1, **({} if error is None else {"error": dict(error)})} + + +def azure_rate_limit() -> Mapping[str, JsonValue]: + return { + "type": "too_many_requests", + "code": "rate_limit_exceeded", + "headers": {"x-ms-fe-error": "true"}, + "message": RATE_LIMIT_MESSAGE, + "param": None, + } + + +def failed(identity: str, code: str, message: str) -> Mapping[str, JsonValue]: + return { + "type": "response.failed", + "sequence_number": 2, + "response": response_object(identity, "failed", error={"code": code, "message": message}), + } + + +def delta(identity: str, text: str) -> Mapping[str, JsonValue]: + return { + "type": "response.output_text.delta", + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": text, + } + + +def completed(identity: str, text: str) -> Mapping[str, JsonValue]: + message: Final = { + "type": "message", + "id": f"msg_{identity}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + usage: Final = { + "input_tokens": 11, + "output_tokens": 4, + "total_tokens": 15, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + } + return { + "type": "response.completed", + "sequence_number": 3, + "response": response_object(identity, "completed", output=[message], usage=usage), + } + + +def rate_limited_stream(identity: str) -> tuple[bytes, ...]: + return ( + frame(created(identity)), + frame(error_event(azure_rate_limit())), + frame(failed(identity, "rate_limit_exceeded", RATE_LIMIT_MESSAGE)), + ) + + +def healthy_stream(identity: str, text: str) -> tuple[bytes, ...]: + return (frame(created(identity)), frame(delta(identity, text)), frame(completed(identity, text))) + + +def serve(stream: tuple[bytes, ...], target: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target.startswith(target), request.target + return Reply(content_type="text/event-stream", chunks=stream) + + return respond + + +def function_tools() -> list[JsonValue]: + return [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Weather for a city", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + }, + } + ] + + +def chat_content(frames: Sequence[Mapping[str, JsonValue]]) -> str: + def deltas() -> Iterator[str]: + for chunk in frames: + for choice in chunk.get("choices") or []: + yield string_value(object_value(object_value(choice)["delta"]).get("content") or "") + + return "".join(deltas()) diff --git a/tests/integration/sdk/test_responses_bridge_stream_errors_sdk.py b/tests/integration/sdk/test_responses_bridge_stream_errors_sdk.py new file mode 100644 index 00000000000..223d6a9eac6 --- /dev/null +++ b/tests/integration/sdk/test_responses_bridge_stream_errors_sdk.py @@ -0,0 +1,137 @@ +import uuid +from typing import Final + +import pytest +from integration._support.responses_stream import ( + AZURE_TARGET, + RATE_LIMIT_MESSAGE, + function_tools, + rate_limited_stream, + serve, +) +from integration._support.wire import Wire, wire_server + +import litellm +from litellm import Router +from litellm.exceptions import MidStreamFallbackError, RateLimitError +from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + +_MODEL: Final = "azure/gpt-6" +_GROUP: Final = "bridged-gpt-6" +_API_KEY: Final = "synthetic-azure-key" +_SENTINEL_PREFIX: Final = "litellm.MidStreamFallbackError: " +_TOOLS: Final = function_tools() + + +def _messages(marker: str) -> list[dict[str, str]]: + return [{"role": "user", "content": marker}] + + +def _router(wire: Wire) -> Router: + return Router( + model_list=[ + {"model_name": _GROUP, "litellm_params": {"model": _MODEL, "api_base": wire.url, "api_key": _API_KEY}} + ], + num_retries=0, + ) + + +def _assert_one_attempt(wire: Wire, marker: str) -> None: + received: Final = wire.drain() + assert len(received) == 1 and marker.encode() in received[0].body, [request.target for request in received] + + +def _assert_wraps_the_provider_exception_once(raised: MidStreamFallbackError, wire: Wire, marker: str) -> None: + inner: Final = raised.original_exception + assert isinstance(inner, RateLimitError), repr(inner) + assert inner.status_code == 429 and RATE_LIMIT_MESSAGE in str(inner), str(inner) + assert raised.status_code == 429, raised.status_code + assert raised.is_pre_first_chunk and raised.generated_content == "", ( + raised.is_pre_first_chunk, + raised.generated_content, + ) + assert str(raised).count(_SENTINEL_PREFIX) == 1, str(raised) + _assert_one_attempt(wire, marker) + + +def _assert_surfaces_the_provider_exception(raised: RateLimitError, wire: Wire, marker: str) -> None: + assert type(raised) is RateLimitError, type(raised) + assert raised.status_code == 429 and RATE_LIMIT_MESSAGE in str(raised), str(raised) + assert _SENTINEL_PREFIX not in str(raised), str(raised) + _assert_one_attempt(wire, marker) + + +def test_sync_completion_stream_in_stream_rate_limit_wraps_the_provider_exception_once() -> None: + marker: Final = uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(f"resp_{marker}"), AZURE_TARGET)) as wire: + response: Final = litellm.completion( + model=_MODEL, + messages=_messages(marker), + tools=_TOOLS, + stream=True, + num_retries=0, + api_base=wire.url, + api_key=_API_KEY, + ) + assert isinstance(response, CustomStreamWrapper), type(response) + with pytest.raises(MidStreamFallbackError) as raised: + for _ in response: + pass + _assert_wraps_the_provider_exception_once(raised.value, wire, marker) + + +async def test_async_completion_stream_in_stream_rate_limit_wraps_the_provider_exception_once() -> None: + marker: Final = uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(f"resp_{marker}"), AZURE_TARGET)) as wire: + response: Final = await litellm.acompletion( + model=_MODEL, + messages=_messages(marker), + tools=_TOOLS, + stream=True, + num_retries=0, + api_base=wire.url, + api_key=_API_KEY, + ) + assert isinstance(response, CustomStreamWrapper), type(response) + with pytest.raises(MidStreamFallbackError) as raised: + async for _ in response: + pass + _assert_wraps_the_provider_exception_once(raised.value, wire, marker) + + +def test_router_sync_stream_in_stream_rate_limit_with_fallbacks_disabled_surfaces_the_provider_exception() -> None: + marker: Final = uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(f"resp_{marker}"), AZURE_TARGET)) as wire: + response: Final = _router(wire).completion( + model=_GROUP, messages=_messages(marker), tools=_TOOLS, stream=True, disable_fallbacks=True + ) + with pytest.raises(RateLimitError) as raised: + for _ in response: + pass + _assert_surfaces_the_provider_exception(raised.value, wire, marker) + + +async def test_router_async_stream_in_stream_rate_limit_without_fallbacks_surfaces_the_provider_exception() -> None: + marker: Final = uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(f"resp_{marker}"), AZURE_TARGET)) as wire: + response: Final = await _router(wire).acompletion( + model=_GROUP, messages=_messages(marker), tools=_TOOLS, stream=True + ) + with pytest.raises(RateLimitError) as raised: + async for _ in response: + pass + _assert_surfaces_the_provider_exception(raised.value, wire, marker) + + +async def test_router_async_stream_in_stream_rate_limit_with_fallbacks_disabled_surfaces_the_provider_exception() -> ( + None +): + marker: Final = uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(f"resp_{marker}"), AZURE_TARGET)) as wire: + response: Final = await _router(wire).acompletion( + model=_GROUP, messages=_messages(marker), tools=_TOOLS, stream=True, disable_fallbacks=True + ) + with pytest.raises(RateLimitError) as raised: + async for _ in response: + pass + _assert_surfaces_the_provider_exception(raised.value, wire, marker) diff --git a/tests/integration/streaming/test_responses_bridge_stream_chaos.py b/tests/integration/streaming/test_responses_bridge_stream_chaos.py new file mode 100644 index 00000000000..ca891584a8e --- /dev/null +++ b/tests/integration/streaming/test_responses_bridge_stream_chaos.py @@ -0,0 +1,368 @@ +import asyncio +import re +import signal +import socket +import threading +import uuid +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal, TypeAlias +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import ( + JSON_OBJECT, + Gateway, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process +from integration._support.responses_stream import ( + AZURE_TARGET, + chat_content, + function_tools, + healthy_stream, + rate_limited_stream, +) +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +_MODEL: Final = "bridged-stream-chaos" +_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_SENTINEL_PREFIX: Final = "litellm.MidStreamFallbackError: " +_RATE_LIMIT_PREFIX: Final = "litellm.RateLimitError: " +_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"anthropic-version": "2023-06-01"}) + +Kind: TypeAlias = Literal["chat_limited", "chat_healthy", "messages_limited", "responses_limited"] +_KINDS: Final[tuple[Kind, ...]] = ("chat_limited", "chat_healthy", "messages_limited", "responses_limited") +_CHAT_KINDS: Final[tuple[Kind, ...]] = ("chat_limited", "chat_healthy") +_LOGGED_KINDS: Final[frozenset[Kind]] = frozenset({"chat_limited", "chat_healthy", "responses_limited"}) + + +@dataclass(frozen=True, slots=True) +class _Call: + kind: Kind + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + call_id: str + + +@dataclass(frozen=True, slots=True) +class _Rig: + port: int + proxy: OwnedProxy + + @property + def gateway(self) -> Gateway: + return self.proxy.gateway + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def _config(port: int, directory: Path) -> Path: + stock: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + router_settings: Final = object_value(stock.get("router_settings") or {}) + path: Final = directory / "bridged-stream-chaos.yaml" + path.write_text( + yaml.safe_dump( + { + **stock, + "model_list": [ + { + "model_name": _MODEL, + "litellm_params": { + "model": "azure/gpt-6", + "api_base": f"http://127.0.0.1:{port}", + "api_key": "synthetic-azure-key", + }, + } + ], + "router_settings": {**router_settings, "num_retries": 0}, + } + ) + ) + return path + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("bridged-stream-chaos") + port: Final = _free_port() + with ( + gateway_from_environment() as shared, + owned_proxy_process(shared, directory, {}, config=_config(port, directory), workers=2) as owned, + ): + yield _Rig(port, owned) + + +def _newest_marker(text: str) -> str | None: + found: Final = _MARKER.findall(text) + return found[-1] if found else None + + +def _respond(request: Request) -> Reply: + assert request.method == "POST" and request.target.startswith(AZURE_TARGET), request.target + marker: Final = _newest_marker(request.body.decode()) + assert marker is not None, request.body + identity: Final = f"resp_{uuid.uuid4().hex}" + if b"chat_healthy" in request.body: + return Reply(content_type="text/event-stream", chunks=healthy_stream(identity, f"answer marker-{marker}")) + return Reply(content_type="text/event-stream", chunks=rate_limited_stream(identity)) + + +def _path(kind: Kind) -> str: + match kind: + case "chat_limited" | "chat_healthy": + return "/v1/chat/completions" + case "messages_limited": + return "/v1/messages" + case "responses_limited": + return "/v1/responses" + + +def _body(call: _Call) -> Mapping[str, JsonValue]: + prompt: Final = f"{call.kind} marker-{call.marker}" + common: Final[Mapping[str, JsonValue]] = { + "model": _MODEL, + "stream": True, + "num_retries": 0, + "cache": {"no-cache": True}, + } + match call.kind: + case "chat_limited" | "chat_healthy": + return {**common, "messages": [{"role": "user", "content": prompt}], "tools": function_tools()} + case "messages_limited": + return { + **common, + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt}], + "tools": [ + { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + } + case "responses_limited": + return {**common, "input": prompt} + + +def _calls(count: int, kinds: tuple[Kind, ...]) -> tuple[_Call, ...]: + return tuple(_Call(kinds[index % len(kinds)], uuid.uuid4().hex) for index in range(count)) + + +async def _send(client: httpx.AsyncClient, key: str, call: _Call) -> _Served: + async with client.stream( + "POST", _path(call.kind), json=_body(call), headers={"Authorization": f"Bearer {key}", **_HEADERS} + ) as response: + raw: Final = await response.aread() + return _Served(call, response.status_code, raw.decode(), response.headers["x-litellm-call-id"]) + + +async def _burst( + gateway: Gateway, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, gateway.key, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _data_frames(text: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: {") + ) + + +def _sse_events(text: str) -> tuple[tuple[str, Mapping[str, JsonValue]], ...]: + def parse(block: str) -> tuple[str, Mapping[str, JsonValue]]: + lines: Final = block.splitlines() + event: Final = next(line.removeprefix("event: ") for line in lines if line.startswith("event: ")) + data: Final = next(line.removeprefix("data: ") for line in lines if line.startswith("data: ")) + return event, JSON_OBJECT.validate_json(data) + + return tuple(parse(block) for block in text.strip().split("\n\n") if "event: " in block) + + +def _assert_answered_in_its_own_shape(served: _Served) -> None: + match served.call.kind: + case "chat_healthy": + assert served.status == 200, served.text + assert chat_content(_data_frames(served.text)) == f"answer marker-{served.call.marker}", served.text + case "chat_limited": + assert served.status == 429, served.text + error: Final = object_value(JSON_OBJECT.validate_json(served.text)["error"]) + assert error["type"] == "throttling_error" and str(error["code"]) == "429", error + message: Final = string_value(error["message"]) + assert message.startswith(_RATE_LIMIT_PREFIX) and _SENTINEL_PREFIX not in message, message + case "messages_limited": + assert served.status == 200, served.text + events: Final = _sse_events(served.text) + assert events[0][0] == "message_start" and events[-1][0] == "error", events + frame_error: Final = object_value(events[-1][1]["error"]) + assert frame_error["type"] == "rate_limit_error", frame_error + frame_message: Final = string_value(frame_error["message"]) + assert frame_message.count(_SENTINEL_PREFIX) == 1 and _RATE_LIMIT_PREFIX in frame_message, frame_message + case "responses_limited": + assert served.status == 200, served.text + kinds: Final = [frame["type"] for frame in _data_frames(served.text)] + assert kinds == ["response.created", "response.failed"], served.text + + +def _assert_forwarded(received: tuple[Request, ...], calls: tuple[_Call, ...]) -> None: + posts: Final = tuple(request for request in received if request.method == "POST") + forwarded: Final = sorted(_newest_marker(request.body.decode()) or "" for request in posts) + assert forwarded == sorted(call.marker for call in calls), forwarded + + +def _rows(call_ids: Sequence[str]) -> Sequence[Mapping[str, JsonValue]]: + return eventually( + lambda: read_rows( + 'SELECT litellm_call_id, status FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = ANY(string_to_array(%s, %s))', + (",".join(call_ids), ","), + ), + lambda found: len(found) >= len(call_ids), + seconds=70, + ) + + +def _assert_each_lands_once(served: tuple[_Served, ...]) -> None: + logged: Final = tuple(item for item in served if item.call.kind in _LOGGED_KINDS) + rows: Final = _rows(tuple(item.call_id for item in logged)) + by_call: Final = {string_value(row["litellm_call_id"]): row for row in rows} + assert len(by_call) == len(rows) == len(logged), rows + for item in logged: + expected: Final = "success" if item.call.kind == "chat_healthy" else "failure" + assert by_call[item.call_id]["status"] == expected, (item.call_id, rows) + + +def _worker_pids(log: Path, count: int) -> tuple[int, ...]: + return eventually( + lambda: tuple(int(found.group(1)) for found in _STARTED_WORKER.finditer(log.read_text())), + lambda pids: len(pids) == count, + seconds=30, + ) + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@dataclass(frozen=True, slots=True) +class _Held: + release: threading.Event + markers: SimpleQueue[str] + + def respond(self, request: Request) -> Reply: + marker: Final = _newest_marker(request.body.decode()) + assert marker is not None, request.body + self.markers.put(marker) + assert self.release.wait(timeout=60), "The burst was never released" + return _respond(request) + + +async def _held_burst(gateway: Gateway, calls: tuple[_Call, ...], held: _Held) -> asyncio.Task[tuple[_Served, ...]]: + burst: Final = asyncio.create_task(_burst(gateway, calls, tolerate_transport_errors=True)) + await asyncio.to_thread(eventually, held.markers.qsize, lambda size: size == len(calls), 60) + return burst + + +async def test_mixed_burst_of_bridged_streams_answers_each_call_in_its_own_shape_and_logs_each_once( + rig: _Rig, +) -> None: + calls: Final = _calls(24, _KINDS) + with wire_server(_respond, port=rig.port) as wire: + served: Final = await _burst(rig.gateway, calls) + assert len(served) == 24 + for item in served: + _assert_answered_in_its_own_shape(item) + _assert_forwarded(wire.drain(), calls) + _assert_each_lands_once(served) + + +async def test_worker_sigkill_mid_burst_leaves_the_sibling_answering_the_bridged_streams(rig: _Rig) -> None: + calls: Final = _calls(20, _CHAT_KINDS) + held: Final = _Held(threading.Event(), SimpleQueue()) + with wire_server(held.respond, port=rig.port) as wire: + workers: Final = _worker_pids(rig.proxy.log, 2) + burst: Final = await _held_burst(rig.gateway, calls, held) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + held.release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered_in_its_own_shape(item) + follow_up: Final = _Call("chat_healthy", uuid.uuid4().hex) + (answered,) = await _burst(rig.gateway, (follow_up,)) + _assert_answered_in_its_own_shape(answered) + _assert_forwarded(wire.drain(), (*calls, follow_up)) + _assert_each_lands_once((*served, answered)) + + +@pytest.mark.timeout(4 * graceful_stop_seconds() + 240) +async def test_proxy_sigterm_mid_burst_drains_the_spend_log_queue_and_the_restarted_proxy_serves( + gateway: Gateway, tmp_path: Path +) -> None: + port: Final = _free_port() + config: Final = _config(port, tmp_path) + calls: Final = _calls(20, _CHAT_KINDS) + held: Final = _Held(threading.Event(), SimpleQueue()) + with wire_server(held.respond, port=port) as wire: + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + burst: Final = await _held_burst(owned.gateway, calls, held) + owned.process.terminate() + held.release.set() + served: Final = await burst + await asyncio.to_thread( + eventually, owned.process.poll, lambda code: code is not None, graceful_stop_seconds() + ) + assert len(served) == 20, len(served) + for item in served: + _assert_answered_in_its_own_shape(item) + _assert_forwarded(wire.drain(), calls) + _assert_each_lands_once(served) + follow_up: Final = _Call("chat_healthy", uuid.uuid4().hex) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted: + (answered,) = await _burst(restarted.gateway, (follow_up,)) + _assert_answered_in_its_own_shape(answered) + _assert_forwarded(wire.drain(), (follow_up,)) + _assert_each_lands_once((answered,)) diff --git a/tests/integration/streaming/test_responses_bridge_stream_errors.py b/tests/integration/streaming/test_responses_bridge_stream_errors.py new file mode 100644 index 00000000000..9700658fa70 --- /dev/null +++ b/tests/integration/streaming/test_responses_bridge_stream_errors.py @@ -0,0 +1,460 @@ +import json +import socket +import uuid +from collections.abc import Iterator, Mapping, Sequence +from contextlib import ExitStack +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import pytest +import yaml +from integration._support.client import ( + JSON_OBJECT, + Gateway, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds, owned_proxy +from integration._support.responses_stream import ( + AZURE_TARGET, + OPENAI_TARGET, + RATE_LIMIT_MESSAGE, + azure_rate_limit, + chat_content, + created, + delta, + error_event, + failed, + frame, + function_tools, + healthy_stream, + rate_limited_stream, + serve, +) +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +_RATE_LIMIT_PREFIX: Final = "litellm.RateLimitError: " +_SENTINEL_PREFIX: Final = "litellm.MidStreamFallbackError: " +_PRIMARY: Final = "bridged-primary" +_SPARE: Final = "bridged-spare" +pytestmark: Final = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + + +def chat_body(model: str, marker: str, **extra: JsonValue) -> dict[str, JsonValue]: + return { + "model": model, + "stream": True, + "messages": [{"role": "user", "content": marker}], + "tools": function_tools(), + "num_retries": 0, + "cache": {"no-cache": True}, + **extra, + } + + +def messages_body(model: str, marker: str) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 64, + "stream": True, + "messages": [{"role": "user", "content": marker}], + "tools": [ + { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + } + ], + "num_retries": 0, + "cache": {"no-cache": True}, + } + + +def error_body(response: httpx.Response) -> Mapping[str, JsonValue]: + return object_value(JSON_OBJECT.validate_json(response.content)["error"]) + + +def data_frames(text: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + + +def spend_row(call_id: str) -> Mapping[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT litellm_call_id, status, model_group FROM "LiteLLM_SpendLogs" WHERE litellm_call_id = %s', + (call_id,), + ), + lambda found: len(found) >= 1, + seconds=70, + ) + assert len(rows) == 1, rows + return rows[0] + + +def assert_provider_typed_rate_limit(error: Mapping[str, JsonValue]) -> None: + assert error["type"] == "throttling_error", error + assert str(error["code"]) == "429", error + message: Final = string_value(error["message"]) + assert message.startswith(_RATE_LIMIT_PREFIX) and RATE_LIMIT_MESSAGE in message, message + assert _SENTINEL_PREFIX not in message, message + + +def assert_failed_once(wire: Wire, call_id: str, model: str, attempts: int = 1) -> tuple[Request, ...]: + received: Final = wire.drain() + assert len(received) == attempts, [request.target for request in received] + row: Final = spend_row(call_id) + assert row["status"] == "failure" and row["model_group"] == model, row + return received + + +def test_bridged_azure_in_stream_rate_limit_reaches_the_openai_sdk_as_a_throttling_error(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(identity), AZURE_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, max_retries=0) + with pytest.raises(openai.RateLimitError) as raised: + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": identity}], + tools=function_tools(), + stream=True, + extra_body={"num_retries": 0, "cache": {"no-cache": True}}, + ) + assert raised.value.status_code == 429 + assert_provider_typed_rate_limit(object_value(raised.value.body)) + assert_failed_once(wire, raised.value.response.headers["x-litellm-call-id"], model) + + +async def test_bridged_openai_in_stream_rate_limit_reaches_the_async_openai_sdk_as_a_throttling_error( + gateway: Gateway, +) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(identity), OPENAI_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="openai/gpt-5.3-codex", api_base=wire.url, api_key="synthetic-openai-key") + client: Final = openai.AsyncOpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, max_retries=0) + with pytest.raises(openai.RateLimitError) as raised: + await client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": identity}], + stream=True, + extra_body={"num_retries": 0, "cache": {"no-cache": True}}, + ) + assert raised.value.status_code == 429 + assert_provider_typed_rate_limit(object_value(raised.value.body)) + assert_failed_once(wire, raised.value.response.headers["x-litellm-call-id"], model) + + +def _chat(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", body) + + +@dataclass(frozen=True, slots=True) +class FallbackProxy: + gateway: Gateway + primary_port: int + spare_port: int + + +def _free_ports(count: int) -> tuple[int, ...]: + with ExitStack() as reserved: + sockets: Final = tuple(reserved.enter_context(socket.socket()) for _ in range(count)) + for reserve in sockets: + reserve.bind(("127.0.0.1", 0)) + return tuple(reserve.getsockname()[1] for reserve in sockets) + + +def _fallback_config(directory: Path, primary_port: int, spare_port: int) -> Path: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + deployments: Final = [ + { + "model_name": name, + "litellm_params": { + "model": "azure/gpt-6", + "api_base": f"http://127.0.0.1:{port}", + "api_key": "synthetic-azure-key", + }, + } + for name, port in ((_PRIMARY, primary_port), (_SPARE, spare_port)) + ] + router_settings: Final = {"num_retries": 0, "disable_cooldowns": True, "fallbacks": [{_PRIMARY: [_SPARE]}]} + path: Final = directory / "bridged-fallbacks.yaml" + path.write_text(yaml.safe_dump({**base, "model_list": deployments, "router_settings": router_settings})) + return path + + +@pytest.fixture(scope="module") +def fallback_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[FallbackProxy]: + directory: Final = tmp_path_factory.mktemp("bridged-fallbacks") + primary_port, spare_port = _free_ports(2) + with ( + gateway_from_environment() as shared, + owned_proxy(shared, directory, {}, config=_fallback_config(directory, primary_port, spare_port)) as owned, + ): + yield FallbackProxy(owned, primary_port, spare_port) + + +def test_bridged_in_stream_rate_limit_falls_back_to_the_healthy_deployment(fallback_proxy: FallbackProxy) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with ( + wire_server(serve(rate_limited_stream(identity), AZURE_TARGET), port=fallback_proxy.primary_port) as primary, + wire_server( + serve(healthy_stream(identity, "fallback answer"), AZURE_TARGET), port=fallback_proxy.spare_port + ) as spare, + ): + response: Final = _chat(fallback_proxy.gateway, chat_body(_PRIMARY, identity)) + assert response.status_code == 200, response.text + content: Final = chat_content(data_frames(response.text)) + assert content == "fallback answer", response.text + assert len(primary.drain()) == 1 and len(spare.drain()) == 1 + row: Final = spend_row(response.headers["x-litellm-call-id"]) + assert row["status"] == "success" and row["model_group"] == _SPARE, row + + +def test_bridged_in_stream_rate_limit_whose_fallback_is_also_rate_limited_answers_a_throttling_error( + fallback_proxy: FallbackProxy, +) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with ( + wire_server(serve(rate_limited_stream(identity), AZURE_TARGET), port=fallback_proxy.primary_port) as primary, + wire_server(serve(rate_limited_stream(identity), AZURE_TARGET), port=fallback_proxy.spare_port) as spare, + ): + response: Final = _chat(fallback_proxy.gateway, chat_body(_PRIMARY, identity)) + assert response.status_code == 429, response.text + assert_provider_typed_rate_limit(error_body(response)) + assert len(primary.drain()) == 1 and len(spare.drain()) == 1 + assert spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" + + +def _in_stream_error_status(gateway: Gateway, stream: tuple[bytes, ...]) -> tuple[httpx.Response, str]: + with wire_server(serve(stream, AZURE_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + response: Final = _chat(gateway, chat_body(model, uuid.uuid4().hex)) + assert_failed_once(wire, response.headers["x-litellm-call-id"], model) + return response, model + + +def test_bridged_in_stream_server_error_reaches_the_client_as_the_provider_error(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = ( + frame(created(identity)), + frame(error_event({"type": "server_error", "code": "server_error", "message": "The server had an error"})), + frame(failed(identity, "server_error", "The server had an error")), + ) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 500, response.text + error: Final = error_body(response) + message: Final = string_value(error["message"]) + assert str(error["code"]) == "500", error + assert message.startswith("litellm.APIError: ") and "The server had an error" in message, message + assert _SENTINEL_PREFIX not in message, message + + +def test_bridged_in_stream_invalid_prompt_is_a_bad_request_on_both_legs(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = ( + frame(created(identity)), + frame(error_event({"type": "invalid_request_error", "code": "invalid_prompt", "message": "Invalid prompt"})), + frame(failed(identity, "invalid_prompt", "Invalid prompt")), + ) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 400, response.text + error: Final = error_body(response) + message: Final = string_value(error["message"]) + assert str(error["code"]) == "400", error + assert message.startswith("litellm.BadRequestError: ") and "Invalid prompt" in message, message + assert _SENTINEL_PREFIX not in message, message + + +def test_bridged_error_event_without_an_error_object_is_a_provider_typed_internal_error(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = (frame(created(identity)), frame(error_event(None))) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 500, response.text + error: Final = error_body(response) + message: Final = string_value(error["message"]) + assert str(error["code"]) == "500", error + assert message.startswith("litellm.APIError: ") and "Response API in-stream error" in message, message + assert _SENTINEL_PREFIX not in message, message + + +def test_bridged_error_event_with_a_numeric_code_is_a_throttling_error(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = (frame(created(identity)), frame(error_event({"code": "429", "message": RATE_LIMIT_MESSAGE}))) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 429, response.text + assert_provider_typed_rate_limit(error_body(response)) + + +def test_bridged_response_failed_without_an_error_event_is_a_throttling_error(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = (frame(created(identity)), frame(failed(identity, "rate_limit_exceeded", RATE_LIMIT_MESSAGE))) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 429, response.text + assert_provider_typed_rate_limit(error_body(response)) + + +def test_bridged_rate_limit_after_output_is_a_provider_typed_error_frame_behind_the_text(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + stream: Final = ( + frame(created(identity)), + frame(delta(identity, "Hello")), + frame(delta(identity, " there")), + frame(error_event(azure_rate_limit())), + frame(failed(identity, "rate_limit_exceeded", RATE_LIMIT_MESSAGE)), + ) + response, _ = _in_stream_error_status(gateway, stream) + assert response.status_code == 200, response.text + frames: Final = data_frames(response.text) + content: Final = chat_content(frames) + assert content == "Hello there", response.text + error: Final = object_value(frames[-1]["error"]) + assert str(error["code"]) == "429", error + assert error["type"] == "throttling_error", error + message: Final = string_value(error["message"]) + assert message.startswith(_RATE_LIMIT_PREFIX) and RATE_LIMIT_MESSAGE in message, message + assert _SENTINEL_PREFIX not in message, message + + +def test_bridged_transport_drop_after_response_created_is_a_500_without_the_sentinel(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.target.startswith(AZURE_TARGET), request.target + return Reply( + content_type="text/event-stream", + chunks=(frame(created(identity)), frame(delta(identity, "never sent"))), + abort_after=1, + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + response: Final = _chat(gateway, chat_body(model, identity)) + assert response.status_code == 500, response.text + error: Final = error_body(response) + message: Final = string_value(error["message"]) + assert str(error["code"]) == "500", error + assert "never sent" not in response.text + assert _SENTINEL_PREFIX not in message, message + assert_failed_once(wire, response.headers["x-litellm-call-id"], model) + + +def test_plain_chat_http_rate_limit_is_a_throttling_error_on_both_legs(gateway: Gateway) -> None: + identity: Final = uuid.uuid4().hex + denial: Final = {"error": {"message": "Rate limit reached", "type": "requests", "code": "rate_limit_exceeded"}} + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/chat/completions", request.target + return Reply(status=429, body=json.dumps(denial).encode()) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="openai/gpt-4o-mini", api_base=wire.url, api_key="synthetic-openai-key") + client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=gateway.key, max_retries=0) + with pytest.raises(openai.RateLimitError) as raised: + client.chat.completions.create( + model=model, + messages=[{"role": "user", "content": identity}], + stream=True, + extra_body={"num_retries": 0, "cache": {"no-cache": True}}, + ) + assert raised.value.status_code == 429 + error: Final = object_value(raised.value.body) + assert error["type"] == "throttling_error" and str(error["code"]) == "429", error + message: Final = string_value(error["message"]) + assert message.startswith(_RATE_LIMIT_PREFIX) and "Rate limit reached" in message, message + assert_failed_once(wire, raised.value.response.headers["x-litellm-call-id"], model) + + +def sse_events(text: str) -> tuple[tuple[str, Mapping[str, JsonValue]], ...]: + def parse(block: str) -> tuple[str, Mapping[str, JsonValue]]: + lines: Final = block.splitlines() + event: Final = next(line.removeprefix("event: ") for line in lines if line.startswith("event: ")) + data: Final = next(line.removeprefix("data: ") for line in lines if line.startswith("data: ")) + return event, JSON_OBJECT.validate_json(data) + + return tuple(parse(block) for block in text.strip().split("\n\n") if "event: " in block) + + +def assert_messages_errorframe(events: Sequence[tuple[str, Mapping[str, JsonValue]]]) -> None: + assert events[0][0] == "message_start", events + assert events[-1][0] == "error", events + error: Final = object_value(events[-1][1]["error"]) + assert error["type"] == "rate_limit_error", error + message: Final = string_value(error["message"]) + assert message.startswith(_SENTINEL_PREFIX + _RATE_LIMIT_PREFIX) and RATE_LIMIT_MESSAGE in message, message + assert message.count(_SENTINEL_PREFIX) == 1, message + + +def test_messages_over_the_bridged_stream_carry_the_provider_error_once_in_the_errorframe(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(identity), AZURE_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + response: Final = gateway.request( + "POST", "/v1/messages", messages_body(model, identity), headers={"anthropic-version": "2023-06-01"} + ) + assert response.status_code == 200, response.text + assert_messages_errorframe(sse_events(response.text)) + assert len(wire.drain()) == 1 + + +async def _consume_anthropic_stream(client: anthropic.AsyncAnthropic, model: str, identity: str) -> None: + async with client.messages.stream( + model=model, + max_tokens=64, + messages=[{"role": "user", "content": identity}], + tools=[ + { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}}, + } + ], + extra_body={"num_retries": 0, "cache": {"no-cache": True}}, + ) as stream: + async for _ in stream: + pass + + +async def test_messages_over_the_bridged_stream_raise_the_error_frame_in_the_anthropic_sdk(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(identity), AZURE_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + client: Final = anthropic.AsyncAnthropic( + base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0 + ) + with pytest.raises(anthropic.APIStatusError) as raised: + await _consume_anthropic_stream(client, model, identity) + body: Final = object_value(raised.value.body) + assert_messages_errorframe((("message_start", {}), ("error", body))) + assert len(wire.drain()) == 1 + + +def test_native_responses_stream_forwards_the_failed_response_on_both_legs(gateway: Gateway) -> None: + identity: Final = "resp_" + uuid.uuid4().hex + with wire_server(serve(rate_limited_stream(identity), AZURE_TARGET)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="azure/gpt-6", api_base=wire.url, api_key="synthetic-azure-key") + response: Final = gateway.request( + "POST", + "/v1/responses", + {"model": model, "input": identity, "stream": True, "cache": {"no-cache": True}}, + ) + assert response.status_code == 200, response.text + frames: Final = data_frames(response.text) + assert [frame["type"] for frame in frames] == ["response.created", "response.failed"], response.text + failed: Final = object_value(frames[-1]["response"]) + assert failed["status"] == "failed", failed + assert object_value(failed["error"])["code"] == "rate_limit_exceeded", failed + assert response.text.rstrip().endswith("data: [DONE]"), response.text + assert len(wire.drain()) == 1 + assert spend_row(response.headers["x-litellm-call-id"])["status"] == "failure" diff --git a/tests/unit/litellm_core_utils/test_streaming_handler.py b/tests/unit/litellm_core_utils/test_streaming_handler.py index 3390c96f38b..01fae0958e6 100644 --- a/tests/unit/litellm_core_utils/test_streaming_handler.py +++ b/tests/unit/litellm_core_utils/test_streaming_handler.py @@ -1,7 +1,7 @@ import asyncio import json import time -from typing import Final, Optional +from typing import Final, NoReturn, Optional from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -848,6 +848,48 @@ def test_sync_streaming_rate_limit_triggers_midstream_fallback(logging_obj: Logg assert excinfo.value.generated_content == "" +@pytest.mark.asyncio +async def test_bridged_stream_mid_stream_fallback_error_is_rebuilt_around_the_provider_error(logging_obj: Logging): + """A MidStreamFallbackError raised by an inner stream (the chat-to-Responses bridge consumes a + Responses stream) is raised once around the provider's RateLimitError, so the Router's one-level + unwrap surfaces it, and carries the outer wrapper's bookkeeping: the inner stream counted the + lifecycle event it yielded as its first chunk, while this wrapper's consumer received nothing.""" + from litellm.exceptions import MidStreamFallbackError, RateLimitError + + rate_limit_error: Final = RateLimitError( + message="Your requests to gpt-6.1-sol have exceeded token rate limit.", + llm_provider="azure", + model="gpt-6.1-sol", + ) + inner_error: Final = MidStreamFallbackError( + message=str(rate_limit_error), + model="gpt-6.1-sol", + llm_provider="azure", + original_exception=rate_limit_error, + is_pre_first_chunk=False, + ) + + async def _raise_inner_error(**kwargs: object) -> NoReturn: + raise inner_error + + response: Final = CustomStreamWrapper( + completion_stream=None, + model="gpt-6.1-sol", + logging_obj=logging_obj, + custom_llm_provider="azure", + make_call=_raise_inner_error, + ) + + with pytest.raises(MidStreamFallbackError) as excinfo: + await response.__anext__() + + assert excinfo.value.original_exception is rate_limit_error + assert excinfo.value.status_code == 429 + assert excinfo.value.message == f"litellm.MidStreamFallbackError: {rate_limit_error}" + assert excinfo.value.is_pre_first_chunk is True + assert excinfo.value.generated_content == "" + + def test_sync_streaming_bad_request_not_midstream(logging_obj: Logging): """Ensure __next__ raises BadRequestError (400) directly, not MidStreamFallbackError.