mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(router): keep include_fallback_errors off the provider call on the sync Router path (#45218)
* fix(router): keep include_fallback_errors off the provider call on the sync Router path Router.completion spread include_fallback_errors into litellm.completion, so OpenAI answered 400 Unknown parameter and a fallback walk lost the backup's reply. Only Router._acompletion popped the router-only flags. Both paths now build the provider call through one shared helper that drops them, and a wire-level respx test covers each path. * test(router): pin include_fallback_errors off the wire across the sync, async, and proxy paths
This commit is contained in:
parent
a3b280f69d
commit
a7ee038592
5 changed files with 1143 additions and 17 deletions
|
|
@ -199,6 +199,7 @@ from litellm.router_utils.common_utils import (
|
|||
resolve_model_group_alias,
|
||||
truncate_fallback_error_detail,
|
||||
warn_on_provider_credential_mismatch,
|
||||
without_router_only_kwargs,
|
||||
)
|
||||
from litellm.router_utils.cooldown_cache import CooldownCache
|
||||
from litellm.router_utils.cooldown_handlers import (
|
||||
|
|
@ -2728,13 +2729,15 @@ class Router:
|
|||
if model in self.model_names or not self.has_model_id(model):
|
||||
self.routing_strategy_pre_call_checks(deployment=deployment)
|
||||
|
||||
input_kwargs: Final = {
|
||||
**litellm_params,
|
||||
"messages": messages,
|
||||
"caching": self.cache_responses,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
input_kwargs: Final = without_router_only_kwargs(
|
||||
{
|
||||
**litellm_params,
|
||||
"messages": messages,
|
||||
"caching": self.cache_responses,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
response: Final = litellm.completion(**input_kwargs)
|
||||
verbose_router_logger.info("litellm.completion(model=%s)\x1b[32m 200 OK\x1b[0m", model_name)
|
||||
|
||||
|
|
@ -3880,15 +3883,15 @@ class Router:
|
|||
)
|
||||
self.total_calls[model_name] += 1
|
||||
|
||||
input_kwargs: Final = {
|
||||
**litellm_params,
|
||||
"messages": messages,
|
||||
"caching": self.cache_responses,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
input_kwargs.pop("silent_model", None)
|
||||
input_kwargs.pop("include_fallback_errors", None)
|
||||
input_kwargs: Final = without_router_only_kwargs(
|
||||
{
|
||||
**litellm_params,
|
||||
"messages": messages,
|
||||
"caching": self.cache_responses,
|
||||
"client": model_client,
|
||||
**kwargs,
|
||||
}
|
||||
)
|
||||
|
||||
logging_obj: Final[LiteLLMLogging | None] = kwargs.get("litellm_logging_obj", None)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ import hashlib
|
|||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final
|
||||
from typing import TYPE_CHECKING, Final, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
|
@ -16,6 +16,14 @@ from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_stru
|
|||
from litellm.types.router import CredentialLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
_V = TypeVar("_V")
|
||||
|
||||
ROUTER_ONLY_CALL_KWARGS: Final = frozenset({"silent_model", "include_fallback_errors"})
|
||||
|
||||
|
||||
def without_router_only_kwargs(kwargs: Mapping[str, _V]) -> dict[str, _V]:
|
||||
return {key: value for key, value in kwargs.items() if key not in ROUTER_ONLY_CALL_KWARGS}
|
||||
|
||||
|
||||
def is_proxy_admin_request(request_kwargs: Mapping[str, object] | None) -> bool:
|
||||
if request_kwargs is None:
|
||||
|
|
|
|||
666
tests/integration/routing/test_include_fallback_errors_wire.py
Normal file
666
tests/integration/routing/test_include_fallback_errors_wire.py
Normal file
|
|
@ -0,0 +1,666 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.anthropic_sse import error_body, message_json, message_stream, parse_sse, stream_reply
|
||||
from integration._support.client import Gateway, eventually, gateway_from_environment
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.openai_wire import chat_reply, responses_reply
|
||||
from integration._support.process import graceful_stop_seconds, owned_proxy
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
from litellm.constants import PROXY_CONFIG_RELOAD_INTERVAL_SECONDS
|
||||
|
||||
Endpoint = Literal["chat", "messages", "responses"]
|
||||
|
||||
_OPENAI_MODEL: Final = "gpt-5.4"
|
||||
_ANTHROPIC_MODEL: Final = "claude-haiku-4-5"
|
||||
_API_KEY: Final = "synthetic-fallback-errors-key"
|
||||
_ROUTER_ONLY_KEYS: Final = ("include_fallback_errors", "silent_model")
|
||||
_ANSWER: Final = "answered by"
|
||||
_UNAUTHORIZED_MESSAGE: Final = "scripted 401: the primary key was revoked"
|
||||
_NONCE: Final = re.compile(r"nonce=([0-9a-f]{32})")
|
||||
_JSON: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_ERRORS: Final = TypeAdapter(list[dict[str, JsonValue]])
|
||||
_ITEMS: Final = TypeAdapter(list[JsonValue])
|
||||
_UPSTREAM_URL_PLACEHOLDER: Final = "upstream-url"
|
||||
_NOT_YET_ON_EVERY_WORKER: Final = (
|
||||
"Invalid model name passed in model=",
|
||||
"There are no healthy deployments for this model",
|
||||
)
|
||||
_PROXY_WORKERS: Final = int(os.environ.get("INTEGRATION_PROXY_WORKERS", "1"))
|
||||
_WORKER_SYNC_SECONDS: Final = 0.0 if _PROXY_WORKERS == 1 else PROXY_CONFIG_RELOAD_INTERVAL_SECONDS + 5.0
|
||||
_FRESH_CONNECTION: Final = MappingProxyType({"Connection": "close"})
|
||||
_SPEND_ROW_SECONDS: Final = 70
|
||||
_BURST: Final = 24
|
||||
_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses")
|
||||
_PATHS: Final = MappingProxyType(
|
||||
{"chat": "/v1/chat/completions", "messages": "/v1/messages", "responses": "/v1/responses"}
|
||||
)
|
||||
_PROVIDERS: Final = MappingProxyType(
|
||||
{
|
||||
"chat": f"openai/{_OPENAI_MODEL}",
|
||||
"messages": f"anthropic/{_ANTHROPIC_MODEL}",
|
||||
"responses": f"openai/{_OPENAI_MODEL}",
|
||||
}
|
||||
)
|
||||
_STREAMING: Final = (pytest.param(False, id="non-stream"), pytest.param(True, id="stream"))
|
||||
_NON_BOOLEAN_FLAGS: Final = (
|
||||
pytest.param(1, True, id="int"),
|
||||
pytest.param("", False, id="empty-string"),
|
||||
pytest.param([], False, id="list"),
|
||||
pytest.param("x" * 5120, True, id="five-kilobytes"),
|
||||
)
|
||||
_OPENAI_UNAUTHORIZED: Final = Reply(
|
||||
status=401,
|
||||
body=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": _UNAUTHORIZED_MESSAGE,
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_api_key",
|
||||
}
|
||||
}
|
||||
).encode(),
|
||||
)
|
||||
_ANTHROPIC_UNAUTHORIZED: Final = Reply(status=401, body=error_body(401, _UNAUTHORIZED_MESSAGE))
|
||||
|
||||
_EXPOSED_OPENAI_PRIMARY: Final = "exposed-openai-primary"
|
||||
_EXPOSED_OPENAI_BACKUP: Final = "exposed-openai-backup"
|
||||
_EXPOSED_OPENAI_SERVING: Final = "exposed-openai-serving"
|
||||
_EXPOSED_OPENAI_FLAKY: Final = "exposed-openai-flaky"
|
||||
_EXPOSED_ANTHROPIC_PRIMARY: Final = "exposed-anthropic-primary"
|
||||
_EXPOSED_ANTHROPIC_BACKUP: Final = "exposed-anthropic-backup"
|
||||
|
||||
_CHAT_BRIDGE_INHERITED_LEAKS: Final = frozenset(
|
||||
f"{deployment}:include_fallback_errors" for deployment in (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP)
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
|
||||
|
||||
|
||||
def _prompt(nonce: str, *, outage: bool = False) -> str:
|
||||
return f"which deployment answers this? nonce={nonce} outage={int(outage)}"
|
||||
|
||||
|
||||
def _nonce_of(text: str) -> str:
|
||||
found: Final = _NONCE.search(text)
|
||||
assert found is not None, text
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def _outage(deployment: str, raw: str) -> bool:
|
||||
return deployment.endswith("primary") or (deployment.endswith("flaky") and "outage=1" in raw)
|
||||
|
||||
|
||||
def _identity(endpoint: Endpoint, deployment: str, nonce: str) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return f"chatcmpl-{deployment}-{nonce}"
|
||||
case "messages":
|
||||
return f"msg_{deployment}_{nonce}"
|
||||
case "responses":
|
||||
return f"resp_{deployment}_{nonce}"
|
||||
|
||||
|
||||
def _peer(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
deployment, _, route = unquote(request.target).lstrip("/").partition("/")
|
||||
raw: Final = request.body.decode()
|
||||
body: Final = _JSON.validate_json(request.body)
|
||||
nonce: Final = _nonce_of(raw)
|
||||
stream: Final = body.get("stream") is True
|
||||
text: Final = f"{_ANSWER} {deployment}"
|
||||
if route == "v1/messages":
|
||||
if _outage(deployment, raw):
|
||||
return _ANTHROPIC_UNAUTHORIZED
|
||||
identity: Final = _identity("messages", deployment, nonce)
|
||||
if stream:
|
||||
return stream_reply(message_stream(identity, _ANTHROPIC_MODEL, text))
|
||||
return Reply(body=message_json(identity, _ANTHROPIC_MODEL, text))
|
||||
if route.endswith("responses"):
|
||||
if _outage(deployment, raw):
|
||||
return _OPENAI_UNAUTHORIZED
|
||||
return responses_reply(_identity("responses", deployment, nonce), _OPENAI_MODEL, text, stream=stream)
|
||||
assert route == "chat/completions", request.target
|
||||
if _outage(deployment, raw):
|
||||
return _OPENAI_UNAUTHORIZED
|
||||
return chat_reply(_identity("chat", deployment, nonce), _OPENAI_MODEL, text, stream=stream)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Posted:
|
||||
deployment: str
|
||||
route: str
|
||||
raw: str
|
||||
|
||||
def leaked(self) -> tuple[str, ...]:
|
||||
return tuple(key for key in _ROUTER_ONLY_KEYS if key in self.raw)
|
||||
|
||||
|
||||
def _posted_of(request: Request) -> _Posted:
|
||||
deployment, _, route = unquote(request.target).lstrip("/").partition("/")
|
||||
return _Posted(deployment, route, request.body.decode())
|
||||
|
||||
|
||||
def _all_posted(wire: Wire) -> tuple[_Posted, ...]:
|
||||
return tuple(_posted_of(request) for request in wire.drain() if request.method == "POST")
|
||||
|
||||
|
||||
def _posted(wire: Wire, nonce: str) -> tuple[_Posted, ...]:
|
||||
return tuple(item for item in _all_posted(wire) if nonce in item.raw)
|
||||
|
||||
|
||||
def _header(response: httpx.Response, name: str) -> str | None:
|
||||
return response.headers[name] if name in response.headers else None
|
||||
|
||||
|
||||
def _leaks(posted: tuple[_Posted, ...]) -> tuple[str, ...]:
|
||||
return tuple(f"{item.deployment}:{','.join(item.leaked())}" for item in posted if item.leaked())
|
||||
|
||||
|
||||
def _hit(posted: tuple[_Posted, ...]) -> tuple[str, ...]:
|
||||
return tuple(item.deployment for item in posted)
|
||||
|
||||
|
||||
def _body(endpoint: Endpoint, model: str, nonce: str, *, stream: bool, outage: bool = False) -> dict[str, JsonValue]:
|
||||
prompt: Final = _prompt(nonce, outage=outage)
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return {"model": model, "stream": stream, "messages": [{"role": "user", "content": prompt}]}
|
||||
case "messages":
|
||||
return {
|
||||
"model": model,
|
||||
"stream": stream,
|
||||
"max_tokens": 32,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
}
|
||||
case "responses":
|
||||
return {"model": model, "stream": stream, "input": prompt}
|
||||
|
||||
|
||||
def _output_item_identity(item: JsonValue) -> str:
|
||||
return str(_JSON.validate_python(item)["id"]).removeprefix("msg_")
|
||||
|
||||
|
||||
def _served_id(endpoint: Endpoint, response: httpx.Response, *, stream: bool) -> str:
|
||||
if not stream:
|
||||
body: Final = _JSON.validate_json(response.content)
|
||||
if endpoint == "responses":
|
||||
return _output_item_identity(_ITEMS.validate_python(body["output"])[0])
|
||||
return str(body["id"])
|
||||
events: Final = parse_sse(response.text)
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return str(events[0].data["id"])
|
||||
case "messages":
|
||||
start: Final = next(event for event in events if event.event == "message_start")
|
||||
return str(_JSON.validate_python(start.data["message"])["id"])
|
||||
case "responses":
|
||||
done: Final = next(event for event in events if event.data.get("type") == "response.output_item.done")
|
||||
return _output_item_identity(done.data["item"])
|
||||
|
||||
|
||||
def _settled(text: str) -> bool:
|
||||
return not any(phrase in text for phrase in _NOT_YET_ON_EVERY_WORKER)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Sent:
|
||||
nonce: str
|
||||
response: httpx.Response
|
||||
|
||||
|
||||
def _send(
|
||||
gateway: Gateway, path: str, body: Callable[[str], Mapping[str, JsonValue]], *, key: str | None = None
|
||||
) -> _Sent:
|
||||
def attempt() -> _Sent:
|
||||
nonce: Final = uuid.uuid4().hex
|
||||
return _Sent(nonce, gateway.request("POST", path, body(nonce), key=key, headers=_FRESH_CONNECTION))
|
||||
|
||||
return eventually(attempt, lambda sent: _settled(sent.response.text), seconds=_WORKER_SYNC_SECONDS + 10)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Observed:
|
||||
status: int
|
||||
served_id: str
|
||||
spend_id: str
|
||||
text: str
|
||||
attempted: str | None
|
||||
errors_header: str | None
|
||||
hit: tuple[str, ...]
|
||||
leaks: tuple[str, ...]
|
||||
|
||||
|
||||
def _observe(endpoint: Endpoint, wire: Wire, sent: _Sent, *, stream: bool) -> _Observed:
|
||||
posted: Final = _posted(wire, sent.nonce)
|
||||
response: Final = sent.response
|
||||
return _Observed(
|
||||
status=response.status_code,
|
||||
served_id=_served_id(endpoint, response, stream=stream) if response.status_code == 200 else "",
|
||||
spend_id=_spend_id(response) if response.status_code == 200 and not stream else "",
|
||||
text=response.text,
|
||||
attempted=_header(response, "x-litellm-attempted-fallbacks"),
|
||||
errors_header=_header(response, "x-litellm-fallback-errors"),
|
||||
hit=_hit(posted),
|
||||
leaks=_leaks(posted),
|
||||
)
|
||||
|
||||
|
||||
def _spend_id(response: httpx.Response) -> str:
|
||||
return str(_JSON.validate_json(response.content)["id"])
|
||||
|
||||
|
||||
def _spend_row_lands(spend_id: str) -> None:
|
||||
eventually(
|
||||
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (spend_id,)),
|
||||
lambda rows: len(rows) == 1,
|
||||
seconds=_SPEND_ROW_SECONDS,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Registered:
|
||||
gateway: Gateway
|
||||
wire: Wire
|
||||
primary: Mapping[Endpoint, str]
|
||||
backup: Mapping[Endpoint, str]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def registered() -> Iterator[_Registered]:
|
||||
with gateway_from_environment() as gateway, wire_server(_peer) as wire, gateway.scenario() as scenario:
|
||||
primary: Final[dict[Endpoint, str]] = {
|
||||
endpoint: scenario.model(model=_PROVIDERS[endpoint], api_key=_API_KEY, api_base=f"{wire.url}/primary")
|
||||
for endpoint in _ENDPOINTS
|
||||
}
|
||||
backup: Final[dict[Endpoint, str]] = {
|
||||
endpoint: scenario.model(model=_PROVIDERS[endpoint], api_key=_API_KEY, api_base=f"{wire.url}/backup")
|
||||
for endpoint in _ENDPOINTS
|
||||
}
|
||||
yield _Registered(gateway, wire, MappingProxyType(primary), MappingProxyType(backup))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("endpoint", _ENDPOINTS)
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_default_gateway_keeps_the_flag_off_the_wire_on_a_fallback(
|
||||
registered: _Registered, endpoint: Endpoint, stream: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
registered.gateway,
|
||||
_PATHS[endpoint],
|
||||
lambda nonce: {
|
||||
**_body(endpoint, registered.primary[endpoint], nonce, stream=stream),
|
||||
"fallbacks": [registered.backup[endpoint]],
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe(endpoint, registered.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert f"{_ANSWER} backup" in observed.text, observed
|
||||
assert observed.served_id == _identity(endpoint, "backup", sent.nonce), observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == ("primary", "backup"), observed
|
||||
assert observed.errors_header is None, observed
|
||||
if not stream:
|
||||
_spend_row_lands(observed.spend_id)
|
||||
if endpoint == "chat" and not stream:
|
||||
assert observed.attempted == "1", observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS)
|
||||
def test_default_gateway_keeps_a_non_boolean_flag_off_the_wire(
|
||||
registered: _Registered, value: JsonValue, errors_reported: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
registered.gateway,
|
||||
_PATHS["chat"],
|
||||
lambda nonce: {
|
||||
**_body("chat", registered.primary["chat"], nonce, stream=False),
|
||||
"fallbacks": [registered.backup["chat"]],
|
||||
"include_fallback_errors": value,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("chat", registered.wire, sent, stream=False)
|
||||
assert observed.status == 200, observed
|
||||
assert f"{_ANSWER} backup" in observed.text, observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == ("primary", "backup"), observed
|
||||
assert observed.errors_header is None, observed
|
||||
|
||||
|
||||
def test_default_gateway_keeps_a_duplicated_raw_flag_off_the_wire(registered: _Registered) -> None:
|
||||
def raw_body(nonce: str) -> str:
|
||||
body: Final = {
|
||||
**_body("chat", registered.primary["chat"], nonce, stream=False),
|
||||
"fallbacks": [registered.backup["chat"]],
|
||||
}
|
||||
return json.dumps(body)[:-1] + ', "include_fallback_errors": true, "include_fallback_errors": true}'
|
||||
|
||||
def attempt() -> _Sent:
|
||||
nonce: Final = uuid.uuid4().hex
|
||||
response: Final = registered.gateway.client.post(
|
||||
_PATHS["chat"],
|
||||
content=raw_body(nonce),
|
||||
headers={
|
||||
"Authorization": f"Bearer {registered.gateway.key}",
|
||||
"Content-Type": "application/json",
|
||||
**_FRESH_CONNECTION,
|
||||
},
|
||||
)
|
||||
return _Sent(nonce, response)
|
||||
|
||||
sent: Final = eventually(attempt, lambda s: _settled(s.response.text), seconds=_WORKER_SYNC_SECONDS + 10)
|
||||
observed: Final = _observe("chat", registered.wire, sent, stream=False)
|
||||
assert observed.status == 200, observed
|
||||
assert observed.served_id == _identity("chat", "backup", sent.nonce), observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == ("primary", "backup"), observed
|
||||
assert observed.errors_header is None, observed
|
||||
|
||||
|
||||
def test_default_gateway_rejects_an_unauthenticated_flagged_request_before_the_wire(registered: _Registered) -> None:
|
||||
nonce: Final = uuid.uuid4().hex
|
||||
response: Final = registered.gateway.request(
|
||||
"POST",
|
||||
_PATHS["chat"],
|
||||
{
|
||||
**_body("chat", registered.primary["chat"], nonce, stream=False),
|
||||
"fallbacks": [registered.backup["chat"]],
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
key="sk-not-a-key",
|
||||
headers=_FRESH_CONNECTION,
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert _posted(registered.wire, nonce) == ()
|
||||
|
||||
|
||||
def _exposed_deployment(name: str, provider: str) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model_name": name,
|
||||
"litellm_params": {
|
||||
"model": provider,
|
||||
"api_base": f"{_UPSTREAM_URL_PLACEHOLDER}/{name}",
|
||||
"api_key": _API_KEY,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _exposed_config(wire: Wire, directory: Path) -> Path:
|
||||
config: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
config["model_list"] = [
|
||||
_exposed_deployment(_EXPOSED_OPENAI_PRIMARY, _PROVIDERS["chat"]),
|
||||
_exposed_deployment(_EXPOSED_OPENAI_BACKUP, _PROVIDERS["chat"]),
|
||||
_exposed_deployment(_EXPOSED_OPENAI_SERVING, _PROVIDERS["chat"]),
|
||||
_exposed_deployment(_EXPOSED_OPENAI_FLAKY, _PROVIDERS["chat"]),
|
||||
_exposed_deployment(_EXPOSED_ANTHROPIC_PRIMARY, _PROVIDERS["messages"]),
|
||||
_exposed_deployment(_EXPOSED_ANTHROPIC_BACKUP, _PROVIDERS["messages"]),
|
||||
]
|
||||
config["general_settings"] = {
|
||||
**_JSON.validate_python(config["general_settings"]),
|
||||
"expose_fallback_errors_to_caller": True,
|
||||
}
|
||||
config["router_settings"] = {
|
||||
"num_retries": 0,
|
||||
"disable_cooldowns": True,
|
||||
"fallbacks": [
|
||||
{_EXPOSED_OPENAI_PRIMARY: [_EXPOSED_OPENAI_BACKUP]},
|
||||
{_EXPOSED_OPENAI_FLAKY: [_EXPOSED_OPENAI_BACKUP]},
|
||||
{_EXPOSED_ANTHROPIC_PRIMARY: [_EXPOSED_ANTHROPIC_BACKUP]},
|
||||
],
|
||||
}
|
||||
path: Final = directory / "include-fallback-errors-exposed.yaml"
|
||||
path.write_text(yaml.safe_dump(config).replace(_UPSTREAM_URL_PLACEHOLDER, wire.url))
|
||||
return path
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Exposed:
|
||||
proxy: Gateway
|
||||
wire: Wire
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def exposed(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Exposed]:
|
||||
directory: Final = tmp_path_factory.mktemp("include-fallback-errors-exposed")
|
||||
with gateway_from_environment() as gateway, wire_server(_peer) as wire:
|
||||
with owned_proxy(gateway, directory, {}, config=_exposed_config(wire, directory), workers=2) as proxy:
|
||||
yield _Exposed(proxy, wire)
|
||||
|
||||
|
||||
def _errors_mention_the_outage(errors_header: str | None) -> bool:
|
||||
if errors_header is None:
|
||||
return False
|
||||
errors: Final = _ERRORS.validate_json(errors_header)
|
||||
return len(errors) == 1 and _UNAUTHORIZED_MESSAGE in str(errors[0]["message"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_exposed_chat_fallback_reports_the_errors_without_putting_the_flag_on_the_wire(
|
||||
exposed: _Exposed, stream: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["chat"],
|
||||
lambda nonce: {
|
||||
**_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("chat", exposed.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert f"{_ANSWER} {_EXPOSED_OPENAI_BACKUP}" in observed.text, observed
|
||||
assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed
|
||||
if not stream:
|
||||
assert observed.attempted == "1", observed
|
||||
assert _errors_mention_the_outage(observed.errors_header), observed
|
||||
|
||||
|
||||
def test_exposed_chat_without_a_fallback_keeps_the_flag_off_the_wire_and_the_errors_header_off(
|
||||
exposed: _Exposed,
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["chat"],
|
||||
lambda nonce: {
|
||||
**_body("chat", _EXPOSED_OPENAI_SERVING, nonce, stream=False),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("chat", exposed.wire, sent, stream=False)
|
||||
assert observed.status == 200, observed
|
||||
assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_SERVING, sent.nonce), observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == (_EXPOSED_OPENAI_SERVING,), observed
|
||||
assert observed.attempted == "0", observed
|
||||
assert observed.errors_header is None, observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS)
|
||||
def test_exposed_chat_keeps_a_non_boolean_flag_off_the_wire_and_the_errors_header_follows_its_truthiness(
|
||||
exposed: _Exposed, value: JsonValue, errors_reported: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["chat"],
|
||||
lambda nonce: {
|
||||
**_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=False),
|
||||
"include_fallback_errors": value,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("chat", exposed.wire, sent, stream=False)
|
||||
assert observed.status == 200, observed
|
||||
assert observed.served_id == _identity("chat", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed
|
||||
assert observed.leaks == (), observed
|
||||
assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed
|
||||
assert observed.attempted == "1", observed
|
||||
assert _errors_mention_the_outage(observed.errors_header) is errors_reported, observed
|
||||
|
||||
|
||||
def test_exposed_proxy_rejects_an_unauthenticated_flagged_request_before_the_wire(exposed: _Exposed) -> None:
|
||||
nonce: Final = uuid.uuid4().hex
|
||||
response: Final = exposed.proxy.request(
|
||||
"POST",
|
||||
_PATHS["chat"],
|
||||
{**_body("chat", _EXPOSED_OPENAI_PRIMARY, nonce, stream=False), "include_fallback_errors": True},
|
||||
key="sk-not-a-key",
|
||||
headers=_FRESH_CONNECTION,
|
||||
)
|
||||
assert response.status_code == 401, response.text
|
||||
assert _posted(exposed.wire, nonce) == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_exposed_messages_fallback_keeps_the_flag_off_the_wire(exposed: _Exposed, stream: bool) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["messages"],
|
||||
lambda nonce: {
|
||||
**_body("messages", _EXPOSED_ANTHROPIC_PRIMARY, nonce, stream=stream),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("messages", exposed.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert observed.served_id == _identity("messages", _EXPOSED_ANTHROPIC_BACKUP, sent.nonce), observed
|
||||
assert observed.hit == (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP), observed
|
||||
assert observed.leaks == (), observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_exposed_responses_fallback_keeps_the_flag_off_the_wire(exposed: _Exposed, stream: bool) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["responses"],
|
||||
lambda nonce: {
|
||||
**_body("responses", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("responses", exposed.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert observed.served_id == _identity("responses", _EXPOSED_OPENAI_BACKUP, sent.nonce), observed
|
||||
assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed
|
||||
assert observed.leaks == (), observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_exposed_messages_on_the_openai_deployment_bridges_the_fallback_with_a_clean_wire(
|
||||
exposed: _Exposed, stream: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["messages"],
|
||||
lambda nonce: {
|
||||
**_body("messages", _EXPOSED_OPENAI_PRIMARY, nonce, stream=stream),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("messages", exposed.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert f"{_ANSWER} {_EXPOSED_OPENAI_BACKUP}" in observed.text, observed
|
||||
assert observed.hit == (_EXPOSED_OPENAI_PRIMARY, _EXPOSED_OPENAI_BACKUP), observed
|
||||
assert observed.leaks == (), observed
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_exposed_responses_on_the_anthropic_deployment_falls_back_through_the_chat_bridge_with_no_new_key_on_the_wire(
|
||||
exposed: _Exposed, stream: bool
|
||||
) -> None:
|
||||
sent: Final = _send(
|
||||
exposed.proxy,
|
||||
_PATHS["responses"],
|
||||
lambda nonce: {
|
||||
**_body("responses", _EXPOSED_ANTHROPIC_PRIMARY, nonce, stream=stream),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
)
|
||||
observed: Final = _observe("responses", exposed.wire, sent, stream=stream)
|
||||
assert observed.status == 200, observed
|
||||
assert f"{_ANSWER} {_EXPOSED_ANTHROPIC_BACKUP}" in observed.text, observed
|
||||
assert observed.hit == (_EXPOSED_ANTHROPIC_PRIMARY, _EXPOSED_ANTHROPIC_BACKUP), observed
|
||||
assert set(observed.leaks) <= _CHAT_BRIDGE_INHERITED_LEAKS, observed
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _BurstRequest:
|
||||
endpoint: Endpoint
|
||||
stream: bool
|
||||
outage: bool
|
||||
nonce: str
|
||||
|
||||
|
||||
def _burst_plan() -> tuple[_BurstRequest, ...]:
|
||||
shapes: Final[tuple[tuple[Endpoint, bool], ...]] = (("chat", False), ("chat", True), ("responses", False))
|
||||
return tuple(
|
||||
_BurstRequest(endpoint, stream, index % 2 == 1, uuid.uuid4().hex)
|
||||
for index, (endpoint, stream) in enumerate(shapes * (_BURST // len(shapes)))
|
||||
)
|
||||
|
||||
|
||||
def _fire(proxy: Gateway, request: _BurstRequest) -> httpx.Response:
|
||||
return proxy.request(
|
||||
"POST",
|
||||
_PATHS[request.endpoint],
|
||||
{
|
||||
**_body(
|
||||
request.endpoint, _EXPOSED_OPENAI_FLAKY, request.nonce, stream=request.stream, outage=request.outage
|
||||
),
|
||||
"include_fallback_errors": True,
|
||||
},
|
||||
headers=_FRESH_CONNECTION,
|
||||
)
|
||||
|
||||
|
||||
def _expected_hit(request: _BurstRequest) -> tuple[str, ...]:
|
||||
if request.outage:
|
||||
return (_EXPOSED_OPENAI_FLAKY, _EXPOSED_OPENAI_BACKUP)
|
||||
return (_EXPOSED_OPENAI_FLAKY,)
|
||||
|
||||
|
||||
def _expected_served_id(request: _BurstRequest) -> str:
|
||||
deployment: Final = _EXPOSED_OPENAI_BACKUP if request.outage else _EXPOSED_OPENAI_FLAKY
|
||||
return _identity(request.endpoint, deployment, request.nonce)
|
||||
|
||||
|
||||
def _check_burst_row(request: _BurstRequest, response: httpx.Response, posted: tuple[_Posted, ...]) -> None:
|
||||
own: Final = tuple(item for item in posted if request.nonce in item.raw)
|
||||
assert response.status_code == 200, (request, response.text)
|
||||
assert _served_id(request.endpoint, response, stream=request.stream) == _expected_served_id(request), request
|
||||
assert _hit(own) == _expected_hit(request), (request, own)
|
||||
assert _leaks(own) == (), (request, own)
|
||||
if request.endpoint == "chat" and not request.stream:
|
||||
assert _header(response, "x-litellm-attempted-fallbacks") == ("1" if request.outage else "0"), request
|
||||
assert _errors_mention_the_outage(_header(response, "x-litellm-fallback-errors")) is request.outage, request
|
||||
|
||||
|
||||
def test_exposed_burst_with_scripted_outages_lands_every_prompt_once_per_hop_with_a_clean_wire(
|
||||
exposed: _Exposed,
|
||||
) -> None:
|
||||
plan: Final = _burst_plan()
|
||||
with ThreadPoolExecutor(max_workers=_BURST) as pool:
|
||||
responses: Final = tuple(pool.map(partial(_fire, exposed.proxy), plan))
|
||||
posted: Final = _all_posted(exposed.wire)
|
||||
for request, response in zip(plan, responses, strict=True):
|
||||
_check_burst_row(request, response, posted)
|
||||
|
|
@ -0,0 +1,364 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import threading
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from functools import partial
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import eventually
|
||||
from integration._support.openai_wire import chat_reply
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm import CustomStreamWrapper, Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
|
||||
from litellm.types.utils import Choices, ModelResponse, ModelResponseStream
|
||||
|
||||
_MODEL: Final = "gpt-5.4"
|
||||
_API_KEY: Final = "synthetic-fallback-errors-key"
|
||||
_ROUTER_ONLY_KEYS: Final = ("include_fallback_errors", "silent_model")
|
||||
_ANSWER: Final = "answered by"
|
||||
_UNAUTHORIZED_MESSAGE: Final = "scripted 401: the primary key was revoked"
|
||||
_NONCE: Final = re.compile(r"nonce=([0-9a-f]{32})")
|
||||
_JSON: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_OBJECT: Final = TypeAdapter(dict[str, object])
|
||||
_ERRORS: Final = TypeAdapter(list[dict[str, JsonValue]])
|
||||
_BURST: Final = 12
|
||||
_CALLBACK_WINDOW: Final = 15.0
|
||||
_STREAMING: Final = (pytest.param(False, id="non-stream"), pytest.param(True, id="stream"))
|
||||
_NON_BOOLEAN_FLAGS: Final = (
|
||||
pytest.param(1, True, id="int"),
|
||||
pytest.param("", False, id="empty-string"),
|
||||
pytest.param([], False, id="list"),
|
||||
pytest.param("x" * 5120, True, id="five-kilobytes"),
|
||||
pytest.param(False, False, id="false"),
|
||||
)
|
||||
_UNAUTHORIZED: Final = Reply(
|
||||
status=401,
|
||||
body=json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": _UNAUTHORIZED_MESSAGE,
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": "invalid_api_key",
|
||||
}
|
||||
}
|
||||
).encode(),
|
||||
)
|
||||
|
||||
|
||||
def _prompt(nonce: str) -> str:
|
||||
return f"which deployment answers this? nonce={nonce}"
|
||||
|
||||
|
||||
def _nonce_of(body: Mapping[str, JsonValue]) -> str:
|
||||
found: Final = _NONCE.search(json.dumps(body))
|
||||
assert found is not None, body
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def _peer(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
deployment, _, route = request.target.lstrip("/").partition("/")
|
||||
assert route == "chat/completions", request.target
|
||||
if deployment == "primary":
|
||||
return _UNAUTHORIZED
|
||||
body: Final = _JSON.validate_json(request.body)
|
||||
identity: Final = f"chatcmpl-{deployment}-{_nonce_of(body)}"
|
||||
return chat_reply(identity, _MODEL, f"{_ANSWER} {deployment}", stream=body.get("stream") is True)
|
||||
|
||||
|
||||
def _deployment(name: str, wire: Wire, **extra: JsonValue) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model_name": name,
|
||||
"litellm_params": {
|
||||
"model": f"openai/{_MODEL}",
|
||||
"api_base": f"{wire.url}/{name}",
|
||||
"api_key": _API_KEY,
|
||||
**extra,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _router(wire: Wire, *, cache_responses: bool = False) -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
_deployment("primary", wire),
|
||||
_deployment("backup", wire),
|
||||
_deployment("serving", wire),
|
||||
_deployment("mirrored", wire, silent_model="shadow"),
|
||||
_deployment("shadow", wire),
|
||||
],
|
||||
fallbacks=[{"primary": ["backup"]}],
|
||||
num_retries=0,
|
||||
disable_cooldowns=True,
|
||||
cache_responses=cache_responses,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Posted:
|
||||
deployment: str
|
||||
raw: str
|
||||
|
||||
def leaked(self) -> tuple[str, ...]:
|
||||
return tuple(key for key in _ROUTER_ONLY_KEYS if key in self.raw)
|
||||
|
||||
def nonce(self) -> str:
|
||||
found: Final = _NONCE.search(self.raw)
|
||||
assert found is not None, self.raw
|
||||
return found.group(1)
|
||||
|
||||
|
||||
def _posted(wire: Wire) -> tuple[_Posted, ...]:
|
||||
return tuple(
|
||||
_Posted(request.target.lstrip("/").partition("/")[0], request.body.decode())
|
||||
for request in wire.drain()
|
||||
if request.method == "POST"
|
||||
)
|
||||
|
||||
|
||||
def _leaks(posted: tuple[_Posted, ...]) -> tuple[str, ...]:
|
||||
return tuple(f"{item.deployment}:{','.join(item.leaked())}" for item in posted if item.leaked())
|
||||
|
||||
|
||||
def _hit(posted: tuple[_Posted, ...]) -> tuple[str, ...]:
|
||||
return tuple(item.deployment for item in posted)
|
||||
|
||||
|
||||
def _nonces_at(posted: tuple[_Posted, ...], deployment: str) -> tuple[str, ...]:
|
||||
return tuple(sorted(item.nonce() for item in posted if item.deployment == deployment))
|
||||
|
||||
|
||||
def _hidden(response: object) -> Mapping[str, object]:
|
||||
return _OBJECT.validate_python(getattr(response, "_hidden_params", None) or {})
|
||||
|
||||
|
||||
def _headers(response: object) -> Mapping[str, object]:
|
||||
return _OBJECT.validate_python(_hidden(response).get("additional_headers") or {})
|
||||
|
||||
|
||||
def _cache_hit(response: object) -> bool:
|
||||
return _hidden(response).get("cache_hit") is True
|
||||
|
||||
|
||||
def _error_messages(headers: Mapping[str, object]) -> tuple[str, ...]:
|
||||
raw: Final = headers.get("x-litellm-fallback-errors")
|
||||
if raw is None:
|
||||
return ()
|
||||
return tuple(str(error["message"]) for error in _ERRORS.validate_json(str(raw)))
|
||||
|
||||
|
||||
def _delta(chunk: object) -> str:
|
||||
assert isinstance(chunk, ModelResponseStream), chunk
|
||||
return "".join(str(choice.delta.content or "") for choice in chunk.choices)
|
||||
|
||||
|
||||
def _messages(nonce: str) -> list[dict[str, str]]:
|
||||
return [{"role": "user", "content": _prompt(nonce)}]
|
||||
|
||||
|
||||
def _typed_messages(nonce: str) -> list[AllMessageValues]:
|
||||
return [ChatCompletionUserMessage(role="user", content=_prompt(nonce))]
|
||||
|
||||
|
||||
def _content(response: object) -> str:
|
||||
assert isinstance(response, ModelResponse), response
|
||||
choice: Final = response.choices[0]
|
||||
assert isinstance(choice, Choices), choice
|
||||
return str(choice.message.content)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
text: str
|
||||
headers: Mapping[str, object]
|
||||
cache_hit: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Outcome:
|
||||
text: str
|
||||
attempted: object
|
||||
errors: tuple[str, ...]
|
||||
hit: tuple[str, ...]
|
||||
leaks: tuple[str, ...]
|
||||
|
||||
|
||||
def _outcome(wire: Wire, served: _Served) -> _Outcome:
|
||||
posted: Final = _posted(wire)
|
||||
return _Outcome(
|
||||
text=served.text,
|
||||
attempted=served.headers.get("x-litellm-attempted-fallbacks"),
|
||||
errors=_error_messages(served.headers),
|
||||
hit=_hit(posted),
|
||||
leaks=_leaks(posted),
|
||||
)
|
||||
|
||||
|
||||
def _complete(router: Router, model: str, *, stream: bool, nonce: str | None = None, **request: object) -> _Served:
|
||||
response: Final = router.completion(
|
||||
model=model, messages=_messages(nonce or uuid.uuid4().hex), stream=stream, **request
|
||||
)
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
return _Served("".join(_delta(chunk) for chunk in response), _headers(response), _cache_hit(response))
|
||||
return _Served(_content(response), _headers(response), _cache_hit(response))
|
||||
|
||||
|
||||
async def _acomplete(router: Router, model: str, *, stream: bool, **request: object) -> _Served:
|
||||
response: Final = await router.acompletion(
|
||||
model=model, messages=_typed_messages(uuid.uuid4().hex), stream=stream, **request
|
||||
)
|
||||
if isinstance(response, CustomStreamWrapper):
|
||||
parts: Final = [_delta(chunk) async for chunk in response]
|
||||
return _Served("".join(parts), _headers(response), _cache_hit(response))
|
||||
return _Served(_content(response), _headers(response), _cache_hit(response))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_sync_fallback_keeps_the_flag_off_the_wire_and_reports_the_errors(stream: bool) -> None:
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
outcome: Final = _outcome(wire, _complete(router, "primary", stream=stream, include_fallback_errors=True))
|
||||
assert outcome.leaks == (), outcome
|
||||
assert outcome.hit == ("primary", "backup"), outcome
|
||||
assert outcome.text == f"{_ANSWER} backup", outcome
|
||||
assert outcome.attempted == 1, outcome
|
||||
if not stream:
|
||||
assert len(outcome.errors) == 1 and _UNAUTHORIZED_MESSAGE in outcome.errors[0], outcome
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stream", _STREAMING)
|
||||
def test_sync_matches_the_async_twin(stream: bool) -> None:
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
twin: Final = _outcome(
|
||||
wire, asyncio.run(_acomplete(router, "primary", stream=stream, include_fallback_errors=True))
|
||||
)
|
||||
observed: Final = _outcome(wire, _complete(router, "primary", stream=stream, include_fallback_errors=True))
|
||||
assert observed == twin, (observed, twin)
|
||||
assert twin.leaks == (), twin
|
||||
assert twin.hit == ("primary", "backup"), twin
|
||||
|
||||
|
||||
def test_without_a_fallback_the_flag_still_stays_off_the_wire() -> None:
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
outcome: Final = _outcome(wire, _complete(router, "serving", stream=False, include_fallback_errors=True))
|
||||
assert outcome.leaks == (), outcome
|
||||
assert outcome.hit == ("serving",), outcome
|
||||
assert outcome.text == f"{_ANSWER} serving", outcome
|
||||
assert outcome.attempted == 0, outcome
|
||||
assert outcome.errors == (), outcome
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("value", "errors_reported"), _NON_BOOLEAN_FLAGS)
|
||||
def test_a_non_boolean_flag_stays_off_the_wire_and_the_errors_header_follows_its_truthiness(
|
||||
value: object, errors_reported: bool
|
||||
) -> None:
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
outcome: Final = _outcome(wire, _complete(router, "primary", stream=False, include_fallback_errors=value))
|
||||
assert outcome.leaks == (), outcome
|
||||
assert outcome.hit == ("primary", "backup"), outcome
|
||||
assert outcome.text == f"{_ANSWER} backup", outcome
|
||||
assert outcome.attempted == 1, outcome
|
||||
assert (len(outcome.errors) == 1) is errors_reported, outcome
|
||||
|
||||
|
||||
class _Recorder(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.mentions_the_flag: bool | None = None
|
||||
self.fired = threading.Event()
|
||||
|
||||
def _record(self, kwargs: Mapping[str, object]) -> None:
|
||||
self.mentions_the_flag = "include_fallback_errors" in json.dumps(kwargs, default=str)
|
||||
self.fired.set()
|
||||
|
||||
def log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._record(kwargs)
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
|
||||
) -> None:
|
||||
self._record(kwargs)
|
||||
|
||||
|
||||
async def _acomplete_and_await_the_logger(router: Router, recorder: _Recorder) -> _Served:
|
||||
served: Final = await _acomplete(router, "serving", stream=False, include_fallback_errors=True)
|
||||
assert await asyncio.get_running_loop().run_in_executor(None, recorder.fired.wait, _CALLBACK_WINDOW)
|
||||
return served
|
||||
|
||||
|
||||
def test_sync_logger_kwargs_carry_the_flag_exactly_as_the_async_ones_do(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
sync_recorder: Final = _Recorder()
|
||||
async_recorder: Final = _Recorder()
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
monkeypatch.setattr(litellm, "callbacks", [sync_recorder])
|
||||
_complete(router, "serving", stream=False, include_fallback_errors=True)
|
||||
assert sync_recorder.fired.wait(_CALLBACK_WINDOW)
|
||||
monkeypatch.setattr(litellm, "callbacks", [async_recorder])
|
||||
asyncio.run(_acomplete_and_await_the_logger(router, async_recorder))
|
||||
posted: Final = _posted(wire)
|
||||
assert _leaks(posted) == (), posted
|
||||
assert (sync_recorder.mentions_the_flag, async_recorder.mentions_the_flag) == (False, False)
|
||||
|
||||
|
||||
def test_silent_model_shadow_traffic_carries_neither_router_only_key() -> None:
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
served: Final = _complete(router, "mirrored", stream=False, include_fallback_errors=True)
|
||||
eventually(wire.received.qsize, lambda count: count >= 2, seconds=20)
|
||||
posted: Final = _posted(wire)
|
||||
assert served.text == f"{_ANSWER} mirrored", served
|
||||
assert tuple(sorted(_hit(posted))) == ("mirrored", "shadow"), posted
|
||||
assert _leaks(posted) == (), posted
|
||||
|
||||
|
||||
def test_a_cache_hit_repeats_the_answer_without_a_wire_request(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "cache", None)
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire, cache_responses=True)
|
||||
nonce: Final = uuid.uuid4().hex
|
||||
first: Final = _complete(router, "serving", stream=False, nonce=nonce, include_fallback_errors=True)
|
||||
posted: Final = _posted(wire)
|
||||
second: Final = _complete(router, "serving", stream=False, nonce=nonce, include_fallback_errors=True)
|
||||
again: Final = _posted(wire)
|
||||
assert _hit(posted) == ("serving",), posted
|
||||
assert _leaks(posted) == (), posted
|
||||
assert again == (), again
|
||||
assert second.text == first.text == f"{_ANSWER} serving"
|
||||
assert (first.cache_hit, second.cache_hit) == (False, True), (first, second)
|
||||
|
||||
|
||||
def _flagged(router: Router, nonce: str) -> _Served:
|
||||
return _complete(router, "primary", stream=False, nonce=nonce, include_fallback_errors=True)
|
||||
|
||||
|
||||
def test_a_sync_burst_lands_every_prompt_once_on_each_side_of_the_fallback() -> None:
|
||||
nonces: Final = tuple(uuid.uuid4().hex for _ in range(_BURST))
|
||||
with wire_server(_peer) as wire:
|
||||
router: Final = _router(wire)
|
||||
with ThreadPoolExecutor(max_workers=_BURST) as pool:
|
||||
served: Final = tuple(pool.map(partial(_flagged, router), nonces))
|
||||
posted: Final = _posted(wire)
|
||||
assert tuple(item.text for item in served) == (f"{_ANSWER} backup",) * _BURST, served
|
||||
assert tuple(item.headers.get("x-litellm-attempted-fallbacks") for item in served) == (1,) * _BURST, served
|
||||
assert _leaks(posted) == (), posted
|
||||
assert _nonces_at(posted, "primary") == tuple(sorted(nonces)), posted
|
||||
assert _nonces_at(posted, "backup") == tuple(sorted(nonces)), posted
|
||||
|
|
@ -24593,3 +24593,88 @@ class CompletionCustomHandler(
|
|||
except Exception:
|
||||
print(f"Assertion Error: {traceback.format_exc()}")
|
||||
self.errors.append(traceback.format_exc())
|
||||
|
||||
|
||||
_FALLBACK_WIRE_PRIMARY: Final = "http://primary.wire.test/v1"
|
||||
_FALLBACK_WIRE_BACKUP: Final = "http://backup.wire.test/v1"
|
||||
_FALLBACK_WIRE_BACKUP_REPLY: Final = {
|
||||
"id": "chatcmpl-backup",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "gpt-5.6",
|
||||
"choices": [{"index": 0, "finish_reason": "stop", "message": {"role": "assistant", "content": "pong"}}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
|
||||
|
||||
def _fallback_wire_router() -> Router:
|
||||
return Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "primary",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-5.6",
|
||||
"api_key": "sk-primary",
|
||||
"api_base": _FALLBACK_WIRE_PRIMARY,
|
||||
"max_retries": 0,
|
||||
},
|
||||
},
|
||||
{
|
||||
"model_name": "backup",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-5.6",
|
||||
"api_key": "sk-backup",
|
||||
"api_base": _FALLBACK_WIRE_BACKUP,
|
||||
"max_retries": 0,
|
||||
},
|
||||
},
|
||||
],
|
||||
fallbacks=[{"primary": ["backup"]}],
|
||||
num_retries=0,
|
||||
)
|
||||
|
||||
|
||||
def _mock_fallback_wire(respx_mock: respx.MockRouter) -> tuple[respx.Route, respx.Route]:
|
||||
primary = respx_mock.post(f"{_FALLBACK_WIRE_PRIMARY}/chat/completions").mock(
|
||||
return_value=httpx.Response(500, json={"error": {"message": "primary overloaded", "type": "server_error"}})
|
||||
)
|
||||
backup = respx_mock.post(f"{_FALLBACK_WIRE_BACKUP}/chat/completions").mock(
|
||||
return_value=httpx.Response(200, json=_FALLBACK_WIRE_BACKUP_REPLY)
|
||||
)
|
||||
return primary, backup
|
||||
|
||||
|
||||
def _assert_fallback_errors_reached_the_caller_and_not_the_wire(
|
||||
response: object, primary: respx.Route, backup: respx.Route
|
||||
) -> None:
|
||||
for route in (primary, backup):
|
||||
assert route.called
|
||||
for call in route.calls:
|
||||
assert "include_fallback_errors" not in json.loads(call.request.content)
|
||||
assert isinstance(response, litellm.ModelResponse)
|
||||
assert response.choices[0].message.content == "pong"
|
||||
headers = response._hidden_params["additional_headers"]
|
||||
assert headers["x-litellm-attempted-fallbacks"] == 1
|
||||
errors = json.loads(headers["x-litellm-fallback-errors"])
|
||||
assert len(errors) == 1
|
||||
assert "primary overloaded" in errors[0]["message"]
|
||||
|
||||
|
||||
def test_sync_completion_keeps_include_fallback_errors_off_the_wire_and_returns_the_errors():
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
primary, backup = _mock_fallback_wire(respx_mock)
|
||||
response = _fallback_wire_router().completion(
|
||||
model="primary", messages=[{"role": "user", "content": "hi"}], include_fallback_errors=True
|
||||
)
|
||||
_assert_fallback_errors_reached_the_caller_and_not_the_wire(response, primary, backup)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_keeps_include_fallback_errors_off_the_wire_and_returns_the_errors(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
primary, backup = _mock_fallback_wire(respx_mock)
|
||||
response = await _fallback_wire_router().acompletion(
|
||||
model="primary", messages=[{"role": "user", "content": "hi"}], include_fallback_errors=True
|
||||
)
|
||||
_assert_fallback_errors_reached_the_caller_and_not_the_wire(response, primary, backup)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue