fix(streaming): stop re-wrapping a bridged stream's MidStreamFallbackError (#44989)

* fix(streaming): stop re-wrapping a bridged stream's MidStreamFallbackError

A chat completion served over the Responses API already gets its mid-stream
error wrapped by the Responses iterator. The chat stream wrapper wrapped it a
second time, and the router unwraps one layer, so the client saw the inner
sentinel (message prefixed litellm.MidStreamFallbackError, type null) instead
of the provider's RateLimitError. The chat wrapper now re-raises an already
wrapped error untouched.

* test(streaming): type the bridged stream regression test's locals

* fix(streaming): rebuild a bridged mid-stream error with the outer wrapper's bookkeeping

* test(integration): cover the bridged stream error typing on every surface

Checked-in audit cells for the Responses bridge: chat completions through the OpenAI SDK and httpx, /v1/messages through httpx and the Anthropic SDK, native /v1/responses, the litellm and Router SDK stream paths, and a chaos file with a mixed burst, a worker SIGKILL and a proxy SIGTERM against an owned two-worker proxy. Every cell scripts the provider through a wire server and asserts the caller's body, the upstream's received requests and the spend row by id.

* test(integration): read bridged chat content through one helper

* test(integration): build the bridged fallback config without mutating the loaded yaml

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-07 16:01:03 -07:00 • committed by GitHub
parent 4417bf08ae
commit 3ca3e1c686
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1156 additions and 1 deletions

View file

@ -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 -

View file

@ -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())

View file

@ -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)

View file

@ -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,))

View file

@ -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"

View file

@ -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.