mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
4417bf08ae
commit
3ca3e1c686
6 changed files with 1156 additions and 1 deletions
|
|
@ -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 -
|
||||
|
|
|
|||
128
tests/integration/_support/responses_stream.py
Normal file
128
tests/integration/_support/responses_stream.py
Normal 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())
|
||||
137
tests/integration/sdk/test_responses_bridge_stream_errors_sdk.py
Normal file
137
tests/integration/sdk/test_responses_bridge_stream_errors_sdk.py
Normal 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)
|
||||
|
|
@ -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,))
|
||||
|
|
@ -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"
|
||||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue