This commit is contained in:
mayank-affirm 2026-10-05 16:37:51 +00:00 • committed by GitHub
commit 23015e8d65
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 1400 additions and 11 deletions

View file

@ -852,21 +852,25 @@ def _map_openai_like_exception(
_BEDROCK_MANTLE_CONTEXT_WINDOW_PATTERN: Final = re.compile(r"prompt tokens \((\d+)\) exceed model maximum \((\d+)\)")
_BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE: Final = (
"prompt is too long: your prompt exceeds the model's context window"
)
def _get_bedrock_mantle_context_window_message(error_str: str) -> str | None:
"""
Mantle reports context overflow as a structured validation error rather than
the plain-text patterns Bedrock itself uses, so it needs its own detection and a
message clients recognize as context overflow (litellm/litellm#36546).
Mantle reports context overflow as a validation_error carrying the token counts, or as
OpenAI's context_length_exceeded code, either in a 400 body or in a streamed error event
with no error type. Clients such as Claude Code only treat "prompt is too long" as
overflow (litellm/litellm#36546).
"""
if "invalid_request_error" not in error_str and "validation_error" not in error_str:
return None
match = _BEDROCK_MANTLE_CONTEXT_WINDOW_PATTERN.search(error_str)
if match is None:
return None
prompt_tokens, max_tokens = match.groups()
return f"prompt is too long: {prompt_tokens} tokens > {max_tokens} maximum"
if match is not None:
prompt_tokens, max_tokens = match.groups()
return f"prompt is too long: {prompt_tokens} tokens > {max_tokens} maximum"
if "context_length_exceeded" in error_str:
return _BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE
return None
def _map_bedrock_exception(

View file

@ -19,7 +19,7 @@ from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.exceptions import MidStreamFallbackError
from litellm.exceptions import ContextWindowExceededError, MidStreamFallbackError
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.anthropic import (
AppliedEdit,
@ -63,7 +63,7 @@ def _optional_attr_sequence(obj: object, name: str) -> Sequence[object]:
def _error_status_and_message(exc: Exception) -> tuple[int, str]:
if isinstance(exc, (BaseLLMException, MidStreamFallbackError)):
if isinstance(exc, (BaseLLMException, MidStreamFallbackError, ContextWindowExceededError)):
return exc.status_code, exc.message
return 500, str(exc) or "Upstream stream ended before completion"

View file

@ -0,0 +1,548 @@
import json
import signal
from collections.abc import Iterator, Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Final
from uuid import uuid4
import httpx
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually, gateway_from_environment, object_value
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "openai.gpt-5.6-luna"
_MANTLE_MODEL: Final = f"bedrock_mantle/{_BACKEND}"
_MANTLE_KEY: Final = "synthetic-mantle-bearer"
_OPENAI_KEY: Final = "synthetic-openai-key"
_RESPONSES_PATH: Final = "/openai/v1/responses"
_GENERIC: Final = "prompt is too long: your prompt exceeds the model's context window"
_UPSTREAM_MESSAGE: Final = (
"Your input exceeds the context window of this model. Please adjust your input and try again."
)
_OPENAI_OVERFLOW_MESSAGE: Final = (
"This model's maximum context length is 128000 tokens. However, your messages resulted in 130000 tokens."
)
_OPENAI_OVERFLOW_MARK: Final = "maximum context length is 128000 tokens"
_INVALID_INPUT_MESSAGE: Final = "Invalid 'input': expected a string or array"
_INVALID_PROMPT_MESSAGE: Final = "Invalid prompt: your prompt was flagged as potentially violating our usage policy."
_FALLBACK_TEXT: Final = "fallback answered"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_OVERFLOW_ENVELOPE: Final[dict[str, JsonValue]] = {
"error": {
"code": "context_length_exceeded",
"message": _UPSTREAM_MESSAGE,
"param": "input",
"type": "invalid_request_error",
}
}
_OVERFLOW_BODY: Final = json.dumps(_OVERFLOW_ENVELOPE).encode()
_BAD_INPUT_BODY: Final = json.dumps(
{"error": {"code": None, "message": _INVALID_INPUT_MESSAGE, "param": "input", "type": "invalid_request_error"}}
).encode()
_OPENAI_OVERFLOW_ERROR: Final[dict[str, JsonValue]] = {
"message": _OPENAI_OVERFLOW_MESSAGE,
"type": "invalid_request_error",
"param": "messages",
"code": "context_length_exceeded",
}
_OPENAI_OVERFLOW_BODY: Final = json.dumps({"error": _OPENAI_OVERFLOW_ERROR}).encode()
def _sse(events: tuple[dict[str, JsonValue], ...]) -> tuple[bytes, ...]:
return tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events)
def _response_object(identity: str, status: str, model: str) -> dict[str, JsonValue]:
return {
"id": f"resp_{identity}",
"object": "response",
"created_at": 1789788253,
"status": status,
"model": model,
"output": [],
}
def _failed_frames(identity: str, error: JsonValue, model: str = _BACKEND) -> tuple[bytes, ...]:
return _sse(
(
{
"type": "response.created",
"sequence_number": 0,
"response": _response_object(identity, "in_progress", model),
},
{
"type": "response.failed",
"sequence_number": 1,
"response": {**_response_object(identity, "failed", model), "error": error},
},
)
)
_PASSTHROUGH_FAILED_FRAMES: Final = _failed_frames("passthrough", _OVERFLOW_ENVELOPE["error"])
def _streaming(request: Request) -> bool:
return _JSON_OBJECT.validate_json(request.body).get("stream") is True
def _overflow_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target
if not _streaming(request):
return Reply(status=400, body=_OVERFLOW_BODY)
return Reply(content_type="text/event-stream", chunks=_failed_frames(uuid4().hex, _OVERFLOW_ENVELOPE["error"]))
def _passthrough_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target
if not _streaming(request):
return Reply(status=400, body=_OVERFLOW_BODY)
return Reply(content_type="text/event-stream", chunks=_PASSTHROUGH_FAILED_FRAMES)
def _bad_input_peer(request: Request) -> Reply:
assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target
if not _streaming(request):
return Reply(status=400, body=_BAD_INPUT_BODY)
error: Final[dict[str, JsonValue]] = {"code": "invalid_prompt", "message": _INVALID_PROMPT_MESSAGE}
return Reply(content_type="text/event-stream", chunks=_failed_frames(uuid4().hex, error))
_EMPTY_MODEL_LIST: Final = Reply(body=b'{"object":"list","data":[]}')
def _openai_overflow_peer(request: Request) -> Reply:
if request.method == "GET":
return _EMPTY_MODEL_LIST
if request.target == "/v1/responses" and _streaming(request):
return Reply(
content_type="text/event-stream",
chunks=_failed_frames(uuid4().hex, _OPENAI_OVERFLOW_ERROR, "gpt-4o-mini"),
)
assert request.target in ("/v1/chat/completions", "/v1/responses"), request.target
return Reply(status=400, body=_OPENAI_OVERFLOW_BODY)
def _chat_reply(identity: str, stream: bool) -> Reply:
if not stream:
return Reply(
body=json.dumps(
{
"id": f"chatcmpl-{identity}",
"object": "chat.completion",
"created": 1789788253,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": _FALLBACK_TEXT},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 9, "completion_tokens": 2, "total_tokens": 11},
}
).encode()
)
chunk: Final[dict[str, JsonValue]] = {
"id": f"chatcmpl-{identity}",
"object": "chat.completion.chunk",
"created": 1789788253,
"model": "gpt-4o-mini",
}
frames: Final = (
{
**chunk,
"choices": [{"index": 0, "delta": {"role": "assistant", "content": _FALLBACK_TEXT}, "finish_reason": None}],
},
{**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
)
return Reply(
content_type="text/event-stream",
chunks=tuple(f"data: {json.dumps(frame)}\n\n".encode() for frame in frames) + (b"data: [DONE]\n\n",),
)
def _responses_reply(identity: str, stream: bool) -> Reply:
completed: Final[dict[str, JsonValue]] = {
**_response_object(identity, "completed", "gpt-4o-mini"),
"output": [
{
"type": "message",
"id": f"msg_{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": _FALLBACK_TEXT, "annotations": []}],
}
],
"usage": {"input_tokens": 9, "output_tokens": 2, "total_tokens": 11},
}
if not stream:
return Reply(body=json.dumps(completed).encode())
return Reply(
content_type="text/event-stream",
chunks=_sse(
(
{
"type": "response.created",
"sequence_number": 0,
"response": {**completed, "status": "in_progress", "output": []},
},
{
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": f"msg_{identity}",
"output_index": 0,
"content_index": 0,
"delta": _FALLBACK_TEXT,
},
{"type": "response.completed", "sequence_number": 2, "response": completed},
)
),
)
def _fallback_peer(request: Request) -> Reply:
if request.method == "GET":
return _EMPTY_MODEL_LIST
assert request.headers["authorization"] == f"Bearer {_OPENAI_KEY}", dict(request.headers)
identity: Final = uuid4().hex
if request.target == "/v1/responses":
return _responses_reply(identity, _streaming(request))
assert request.target == "/v1/chat/completions", request.target
return _chat_reply(identity, _streaming(request))
@dataclass(frozen=True, slots=True)
class Rig:
gateway: Gateway
owned: OwnedProxy
fallback: Wire
peers: Mapping[str, Wire]
models: Mapping[str, str]
def _mantle_params(api_base: str) -> dict[str, JsonValue]:
return {"model": _MANTLE_MODEL, "api_base": api_base, "api_key": _MANTLE_KEY}
def _openai_params(api_base: str) -> dict[str, JsonValue]:
return {"model": "openai/gpt-4o-mini", "api_base": api_base + "/v1", "api_key": _OPENAI_KEY}
def _write_config(
root: Path,
name: str,
model_list: list[dict[str, JsonValue]],
router_settings: dict[str, JsonValue],
pass_through_target: str | None,
) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["model_list"] = model_list
config["router_settings"] = {"disable_cooldowns": True, "num_retries": 0, **router_settings}
if pass_through_target is not None:
config["general_settings"]["pass_through_endpoints"] = [
{
"path": "/mantle-passthrough",
"target": pass_through_target,
"headers": {"Authorization": f"Bearer {_MANTLE_KEY}"},
"auth": True,
}
]
path: Final = root / f"{name}.yaml"
path.write_text(yaml.safe_dump(config))
return path
@pytest.fixture(scope="module")
def cwf_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
root: Final = tmp_path_factory.mktemp("mantle-cwf")
rig_id: Final = uuid4().hex[:8]
cwf: Final = f"mantle-cwf-{rig_id}"
fallback_name: Final = f"fallback-{rig_id}"
with (
gateway_from_environment() as gateway,
wire_server(_overflow_peer) as overflow,
wire_server(_fallback_peer) as fallback,
):
config: Final = _write_config(
root,
"cwf",
[
{"model_name": cwf, "litellm_params": _mantle_params(overflow.url)},
{"model_name": fallback_name, "litellm_params": _openai_params(fallback.url)},
],
{"context_window_fallbacks": [{cwf: [fallback_name]}]},
None,
)
with owned_proxy_process(gateway, root, {}, config=config, workers=2) as owned:
yield Rig(owned.gateway, owned, fallback, {"overflow": overflow}, {"cwf": cwf})
@pytest.fixture(scope="module")
def fb_rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Rig]:
root: Final = tmp_path_factory.mktemp("mantle-fb")
rig_id: Final = uuid4().hex[:8]
names: Final = {
"fb": f"mantle-fb-{rig_id}",
"bad": f"mantle-bad-{rig_id}",
"openai_fb": f"openai-fb-{rig_id}",
}
fallback_name: Final = f"fallback-{rig_id}"
with (
gateway_from_environment() as gateway,
wire_server(_overflow_peer) as overflow,
wire_server(_passthrough_peer) as passthrough,
wire_server(_bad_input_peer) as bad,
wire_server(_openai_overflow_peer) as openai_overflow,
wire_server(_fallback_peer) as fallback,
):
config: Final = _write_config(
root,
"fb",
[
{"model_name": names["fb"], "litellm_params": _mantle_params(overflow.url)},
{"model_name": names["bad"], "litellm_params": _mantle_params(bad.url)},
{"model_name": names["openai_fb"], "litellm_params": _openai_params(openai_overflow.url)},
{"model_name": fallback_name, "litellm_params": _openai_params(fallback.url)},
],
{"fallbacks": [{name: [fallback_name]} for name in names.values()]},
passthrough.url + _RESPONSES_PATH,
)
with owned_proxy_process(gateway, root, {}, config=config, workers=2) as owned:
yield Rig(
owned.gateway,
owned,
fallback,
{"overflow": overflow, "passthrough": passthrough, "bad": bad, "openai_overflow": openai_overflow},
names,
)
def _chat(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": stream}
def _messages(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": prompt}], "stream": stream}
def _responses(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "input": prompt, "stream": stream}
def _consumed(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> tuple[httpx.Response, str]:
with gateway.client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"}) as response:
text: Final = b"".join(response.iter_bytes()).decode()
return response, text
def _events(text: str) -> tuple[dict[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 _only_error(text: str) -> dict[str, JsonValue]:
errors: Final = tuple(event for event in _events(text) if event.get("type") == "error")
assert len(errors) == 1, text
return object_value(errors[0]["error"])
def _only_failed(text: str) -> dict[str, JsonValue]:
failed: Final = tuple(event for event in _events(text) if event.get("type") == "response.failed")
assert len(failed) == 1, text
return object_value(object_value(failed[0]["response"])["error"])
def _streamed_text(text: str) -> str:
return "".join(
str(object_value(event["delta"])["text"]) for event in _events(text) if event["type"] == "content_block_delta"
)
def _calls_with(wire: Wire, prompt: str) -> int:
return sum(1 for request in wire.drain() if prompt in request.body.decode())
def _prompt(label: str) -> str:
return f"{label} {uuid4().hex}"
def _assert_fell_back(rig: Rig, peer: str, prompt: str, response_text: str) -> None:
assert _FALLBACK_TEXT in response_text, response_text
assert _calls_with(rig.peers[peer], prompt) == 1
assert _calls_with(rig.fallback, prompt) == 1
def _assert_no_fallback(rig: Rig, peer: str, prompt: str) -> None:
assert _calls_with(rig.peers[peer], prompt) == 1
assert _calls_with(rig.fallback, prompt) == 0
def test_context_window_fallback_fires_for_the_overflow_on_chat_completions(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf chat")
response: Final = cwf_rig.gateway.request("POST", "/v1/chat/completions", _chat(cwf_rig.models["cwf"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(cwf_rig, "overflow", prompt, response.text)
def test_context_window_fallback_fires_for_the_overflow_on_messages(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf messages")
response: Final = cwf_rig.gateway.request("POST", "/v1/messages", _messages(cwf_rig.models["cwf"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(cwf_rig, "overflow", prompt, response.text)
def test_context_window_fallback_fires_for_the_overflow_on_responses(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf responses")
response: Final = cwf_rig.gateway.request("POST", "/v1/responses", _responses(cwf_rig.models["cwf"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(cwf_rig, "overflow", prompt, response.text)
def test_context_window_fallback_does_not_fire_on_streamed_chat_completions(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf chat stream")
response, text = _consumed(cwf_rig.gateway, "/v1/chat/completions", _chat(cwf_rig.models["cwf"], prompt, True))
assert response.status_code == 400, text
assert _GENERIC in text, text
_assert_no_fallback(cwf_rig, "overflow", prompt)
def test_context_window_fallback_is_not_consulted_on_streamed_messages(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf messages stream")
response, text = _consumed(cwf_rig.gateway, "/v1/messages", _messages(cwf_rig.models["cwf"], prompt, True))
assert response.status_code == 200, text
error: Final = _only_error(text)
assert error["type"] == "invalid_request_error", text
assert _GENERIC in str(error["message"]), text
_assert_no_fallback(cwf_rig, "overflow", prompt)
def test_context_window_fallback_does_not_fire_on_streamed_responses(cwf_rig: Rig) -> None:
prompt: Final = _prompt("cwf responses stream")
response, text = _consumed(cwf_rig.gateway, "/v1/responses", _responses(cwf_rig.models["cwf"], prompt, True))
assert response.status_code == 200, text
assert _GENERIC in str(_only_failed(text)["message"]), text
_assert_no_fallback(cwf_rig, "overflow", prompt)
def test_plain_fallback_fires_for_the_overflow_on_chat_completions(fb_rig: Rig) -> None:
prompt: Final = _prompt("fb chat")
response: Final = fb_rig.gateway.request("POST", "/v1/chat/completions", _chat(fb_rig.models["fb"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(fb_rig, "overflow", prompt, response.text)
def test_plain_fallback_fires_for_the_overflow_on_messages(fb_rig: Rig) -> None:
prompt: Final = _prompt("fb messages")
response: Final = fb_rig.gateway.request("POST", "/v1/messages", _messages(fb_rig.models["fb"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(fb_rig, "overflow", prompt, response.text)
def test_plain_fallback_fires_for_the_overflow_on_responses(fb_rig: Rig) -> None:
prompt: Final = _prompt("fb responses")
response: Final = fb_rig.gateway.request("POST", "/v1/responses", _responses(fb_rig.models["fb"], prompt))
assert response.status_code == 200, response.text
_assert_fell_back(fb_rig, "overflow", prompt, response.text)
def test_plain_fallback_does_not_fire_on_streamed_chat_completions(fb_rig: Rig) -> None:
prompt: Final = _prompt("fb chat stream")
response, text = _consumed(fb_rig.gateway, "/v1/chat/completions", _chat(fb_rig.models["fb"], prompt, True))
assert response.status_code == 400, text
assert _GENERIC in text, text
_assert_no_fallback(fb_rig, "overflow", prompt)
def test_streamed_messages_overflow_on_a_plain_fallback_deployment_returns_the_error_event_without_falling_back(
fb_rig: Rig,
) -> None:
prompt: Final = _prompt("fb messages stream")
response, text = _consumed(fb_rig.gateway, "/v1/messages", _messages(fb_rig.models["fb"], prompt, True))
assert response.status_code == 200, text
error: Final = _only_error(text)
assert error["type"] == "invalid_request_error", text
assert _GENERIC in str(error["message"]), text
_assert_no_fallback(fb_rig, "overflow", prompt)
def test_plain_fallback_does_not_fire_on_streamed_responses(fb_rig: Rig) -> None:
prompt: Final = _prompt("fb responses stream")
response, text = _consumed(fb_rig.gateway, "/v1/responses", _responses(fb_rig.models["fb"], prompt, True))
assert response.status_code == 200, text
assert _GENERIC in str(_only_failed(text)["message"]), text
_assert_no_fallback(fb_rig, "overflow", prompt)
def test_streamed_messages_overflow_on_an_openai_fallback_deployment_returns_the_error_event_without_falling_back(
fb_rig: Rig,
) -> None:
prompt: Final = _prompt("openai fb messages stream")
response, text = _consumed(fb_rig.gateway, "/v1/messages", _messages(fb_rig.models["openai_fb"], prompt, True))
assert response.status_code == 200, text
error: Final = _only_error(text)
assert error["type"] == "invalid_request_error", text
assert _OPENAI_OVERFLOW_MARK in str(error["message"]), text
_assert_no_fallback(fb_rig, "openai_overflow", prompt)
def test_streamed_messages_non_overflow_failure_still_falls_back(fb_rig: Rig) -> None:
prompt: Final = _prompt("bad messages stream")
response, text = _consumed(fb_rig.gateway, "/v1/messages", _messages(fb_rig.models["bad"], prompt, True))
assert response.status_code == 200, text
assert _streamed_text(text) == _FALLBACK_TEXT, text
assert not any(event["type"] == "error" for event in _events(text)), text
_assert_fell_back(fb_rig, "bad", prompt, _FALLBACK_TEXT)
def test_pass_through_relays_the_overflow_envelope_verbatim(fb_rig: Rig) -> None:
prompt: Final = _prompt("passthrough")
response: Final = fb_rig.gateway.request("POST", "/mantle-passthrough", {"model": _BACKEND, "input": prompt})
assert response.status_code == 400, response.text
assert _JSON_OBJECT.validate_json(response.content) == _OVERFLOW_ENVELOPE, response.text
assert _calls_with(fb_rig.peers["passthrough"], prompt) == 1
def test_pass_through_relays_the_failed_stream_frames_verbatim(fb_rig: Rig) -> None:
prompt: Final = _prompt("passthrough stream")
body: Final[dict[str, JsonValue]] = {"model": _BACKEND, "input": prompt, "stream": True}
response, text = _consumed(fb_rig.gateway, "/mantle-passthrough", body)
assert response.status_code == 200, text
assert text.encode() == b"".join(_PASSTHROUGH_FAILED_FRAMES), text
assert _calls_with(fb_rig.peers["passthrough"], prompt) == 1
def _worker_pids(rig: Rig) -> tuple[int, ...]:
children: Final = psutil.Process(rig.owned.process.pid).children(recursive=True)
spawned: Final = tuple(child.pid for child in children if "spawn_main" in " ".join(child.cmdline()))
return spawned or tuple(child.pid for child in children)
def _fallback_status(rig: Rig, prompt: str) -> int:
try:
return rig.gateway.request("POST", "/v1/chat/completions", _chat(rig.models["fb"], prompt)).status_code
except httpx.TransportError:
return -1
def test_killing_one_worker_leaves_the_sibling_serving_the_fallback(fb_rig: Rig) -> None:
workers: Final = eventually(lambda: _worker_pids(fb_rig), lambda pids: len(pids) >= 2, seconds=30)
psutil.Process(workers[0]).send_signal(signal.SIGKILL)
eventually(lambda: _fallback_status(fb_rig, _prompt("after kill")), lambda status: status == 200, seconds=30)
responses: Final = tuple(
fb_rig.gateway.request("POST", "/v1/chat/completions", _chat(fb_rig.models["fb"], _prompt("after kill")))
for _ in range(3)
)
for response in responses:
assert response.status_code == 200, response.text
assert _FALLBACK_TEXT in response.text, response.text
assert fb_rig.gateway.request("GET", "/health/liveliness").status_code == 200

View file

@ -0,0 +1,766 @@
import json
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from itertools import product
from typing import Final
from uuid import uuid4
import anthropic
import httpx
import openai
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "openai.gpt-5.6-luna"
_MODEL: Final = f"bedrock_mantle/{_BACKEND}"
_API_KEY: Final = "synthetic-mantle-bearer"
_RESPONSES_PATH: Final = "/openai/v1/responses"
_OPENAI_MODEL: Final = "openai/gpt-4o-mini"
_OPENAI_CHAT_PATH: Final = "/v1/chat/completions"
_OPENAI_RESPONSES_PATH: Final = "/v1/responses"
_OPENAI_API_KEY: Final = "synthetic-openai-key"
_GENERIC: Final = "prompt is too long: your prompt exceeds the model's context window"
_TOO_LONG: Final = "prompt is too long"
_UPSTREAM_MESSAGE: Final = (
"Your input exceeds the context window of this model. Please adjust your input and try again."
)
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_JSON_VALUE: Final = TypeAdapter(JsonValue)
_INVALID_INPUT_MESSAGE: Final = "Invalid 'input': expected a string or array"
_INVALID_PROMPT_MESSAGE: Final = "Invalid prompt: your prompt was flagged as potentially violating our usage policy."
_OPENAI_OVERFLOW_MESSAGE: Final = (
"This model's maximum context length is 128000 tokens. However, your messages resulted in 130000 tokens. "
"Please reduce the length of the messages."
)
_OPENAI_OVERFLOW_MARK: Final = "maximum context length is 128000 tokens"
_LEGACY_PROMPT_TOKENS: Final = 1055489
_LEGACY_MODEL_MAXIMUM: Final = 1050000
_LEGACY_MESSAGE: Final = (
f"prompt tokens ({_LEGACY_PROMPT_TOKENS}) exceed model maximum ({_LEGACY_MODEL_MAXIMUM}) for {_BACKEND}"
)
_HAPPY_TEXT: Final = "mantle overflow audit control"
def _envelope(code: str | None, message: str) -> bytes:
return json.dumps(
{"error": {"code": code, "message": message, "param": "input", "type": "invalid_request_error"}}
).encode()
_OVERFLOW_BODY: Final = _envelope("context_length_exceeded", _UPSTREAM_MESSAGE)
_LEGACY_BODY: Final = _envelope("validation_error", _LEGACY_MESSAGE)
_BAD_INPUT_BODY: Final = _envelope(None, _INVALID_INPUT_MESSAGE)
_OPENAI_OVERFLOW_ERROR: Final[dict[str, JsonValue]] = {
"message": _OPENAI_OVERFLOW_MESSAGE,
"type": "invalid_request_error",
"param": "messages",
"code": "context_length_exceeded",
}
_OPENAI_OVERFLOW_BODY: Final = json.dumps({"error": _OPENAI_OVERFLOW_ERROR}).encode()
_OVERFLOW: Final = Reply(status=400, body=_OVERFLOW_BODY)
_OVERFLOW_STREAM_ERROR: Final[dict[str, JsonValue]] = {"code": "context_length_exceeded", "message": _UPSTREAM_MESSAGE}
def _response_object(identity: str, status: str, text: str | None, model: str = _BACKEND) -> dict[str, JsonValue]:
output: Final[list[JsonValue]] = (
[]
if text is None
else [
{
"type": "message",
"id": f"msg_{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text, "annotations": []}],
}
]
)
return {
"id": f"resp_{identity}",
"object": "response",
"created_at": 1789788253,
"status": status,
"model": model,
"output": output,
"usage": {"input_tokens": 21, "output_tokens": 4, "total_tokens": 25},
}
def _frames(events: tuple[dict[str, JsonValue], ...]) -> tuple[bytes, ...]:
return tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events)
def _failed_stream(error: JsonValue, model: str = _BACKEND) -> Reply:
identity: Final = uuid4().hex
return Reply(
content_type="text/event-stream",
chunks=_frames(
(
{
"type": "response.created",
"sequence_number": 0,
"response": _response_object(identity, "in_progress", None, model),
},
{
"type": "response.failed",
"sequence_number": 1,
"response": {**_response_object(identity, "failed", None, model), "error": error},
},
)
),
)
def _happy_reply(stream: bool) -> Reply:
identity: Final = uuid4().hex
completed: Final = _response_object(identity, "completed", _HAPPY_TEXT)
if not stream:
return Reply(body=json.dumps(completed).encode())
return Reply(
content_type="text/event-stream",
chunks=_frames(
(
{
"type": "response.created",
"sequence_number": 0,
"response": _response_object(identity, "in_progress", None),
},
{
"type": "response.output_text.delta",
"sequence_number": 1,
"item_id": f"msg_{identity}",
"output_index": 0,
"content_index": 0,
"delta": _HAPPY_TEXT,
},
{"type": "response.completed", "sequence_number": 2, "response": completed},
)
),
)
def _mantle_peer(prompt: str, reply: Reply) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target
assert request.headers["authorization"] == f"Bearer {_API_KEY}", dict(request.headers)
body: Final = _JSON_OBJECT.validate_json(request.body)
assert body["model"] == _BACKEND, body
assert prompt in json.dumps(body["input"]), body
return reply
return respond
def _overflowing_mantle_peer(prompt: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target
body: Final = _JSON_OBJECT.validate_json(request.body)
assert prompt in json.dumps(body["input"]), body
return _failed_stream(_OVERFLOW_STREAM_ERROR) if body.get("stream") is True else _OVERFLOW
return respond
def _openai_peer(prompt: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.method == "GET":
return Reply(body=b'{"object":"list","data":[]}')
assert request.headers["authorization"] == f"Bearer {_OPENAI_API_KEY}", dict(request.headers)
assert prompt in request.body.decode(), request.body
body: Final = _JSON_OBJECT.validate_json(request.body)
if request.target == _OPENAI_RESPONSES_PATH and body.get("stream") is True:
return _failed_stream(_OPENAI_OVERFLOW_ERROR, "gpt-4o-mini")
assert request.target in (_OPENAI_CHAT_PATH, _OPENAI_RESPONSES_PATH), request.target
return Reply(status=400, body=_OPENAI_OVERFLOW_BODY)
return respond
def _chat(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "messages": [{"role": "user", "content": prompt}], "stream": stream}
def _messages(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": prompt}], "stream": stream}
def _responses(model: str, prompt: str, stream: bool = False) -> dict[str, JsonValue]:
return {"model": model, "input": prompt, "stream": stream}
def _consumed(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> tuple[httpx.Response, str]:
with gateway.client.stream("POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}"}) as response:
text: Final = b"".join(response.iter_bytes()).decode()
return response, text
def _events(text: str) -> tuple[dict[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 _error_events(text: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(event for event in _events(text) if event.get("type") == "error")
def _only_error(text: str) -> dict[str, JsonValue]:
errors: Final = _error_events(text)
assert len(errors) == 1, text
return object_value(errors[0]["error"])
def _only_failed(text: str) -> dict[str, JsonValue]:
failed: Final = tuple(event for event in _events(text) if event.get("type") == "response.failed")
assert len(failed) == 1, text
return object_value(object_value(failed[0]["response"])["error"])
def _error_object(response: httpx.Response) -> dict[str, JsonValue]:
return object_value(_JSON_OBJECT.validate_json(response.content)["error"])
def _error_message(response: httpx.Response) -> str:
return str(_error_object(response)["message"])
def _only_call(wire: Wire, target: str = _RESPONSES_PATH) -> None:
assert [(request.method, request.target) for request in wire.drain()] == [("POST", target)]
def _spend_rows(call_id: str) -> list[dict[str, JsonValue]]:
return read_rows('SELECT status, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,))
def _single_row(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(lambda: _spend_rows(call_id), lambda values: len(values) == 1, seconds=70)
return rows[0]
def _failure_message(call_id: str) -> str:
row: Final = _single_row(call_id)
assert row["status"] == "failure", row
metadata: Final = row["metadata"]
parsed: Final = _JSON_OBJECT.validate_json(metadata) if isinstance(metadata, str) else object_value(metadata)
return str(object_value(parsed["error_information"])["error_message"])
def _proxy_url(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _openai_client(gateway: Gateway) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{_proxy_url(gateway)}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(timeout=15, trust_env=False),
)
def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
return anthropic.Anthropic(
base_url=_proxy_url(gateway),
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(timeout=15, trust_env=False),
)
def test_chat_completions_overflow_envelope_returns_400_prompt_too_long_and_logs_failure(gateway: Gateway) -> None:
prompt: Final = f"overflow chat {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 400, response.text
assert _GENERIC in _error_message(response), response.text
_only_call(wire)
assert _GENERIC in _failure_message(response.headers["x-litellm-call-id"])
def test_chat_completions_stream_overflow_envelope_returns_400_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow chat stream {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/chat/completions", _chat(model, prompt, stream=True))
assert response.status_code == 400, text
assert _GENERIC in text, text
_only_call(wire)
def test_messages_overflow_envelope_returns_400_invalid_request_and_logs_failure(gateway: Gateway) -> None:
prompt: Final = f"overflow messages {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/messages", _messages(model, prompt))
assert response.status_code == 400, response.text
error: Final = _error_object(response)
assert error["type"] == "invalid_request_error", response.text
assert _GENERIC in str(error["message"]), response.text
_only_call(wire)
assert _GENERIC in _failure_message(response.headers["x-litellm-call-id"])
def test_messages_stream_overflow_envelope_before_the_stream_returns_400_invalid_request(gateway: Gateway) -> None:
prompt: Final = f"overflow messages stream {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 400, text
error: Final = object_value(_JSON_OBJECT.validate_json(text)["error"])
assert error["type"] == "invalid_request_error", text
assert _GENERIC in str(error["message"]), text
_only_call(wire)
def test_responses_overflow_envelope_returns_400_prompt_too_long_and_logs_failure(gateway: Gateway) -> None:
prompt: Final = f"overflow responses {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/responses", _responses(model, prompt))
assert response.status_code == 400, response.text
assert _GENERIC in _error_message(response), response.text
_only_call(wire)
assert _GENERIC in _failure_message(response.headers["x-litellm-call-id"])
def test_responses_stream_overflow_envelope_before_the_stream_returns_400_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow responses stream {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/responses", _responses(model, prompt, stream=True))
assert response.status_code == 400, text
assert _GENERIC in str(object_value(_JSON_OBJECT.validate_json(text)["error"])["message"]), text
_only_call(wire)
def test_openai_sdk_chat_completions_raises_bad_request_saying_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow sdk chat {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
with pytest.raises(openai.BadRequestError) as caught:
_openai_client(gateway).chat.completions.create(model=model, messages=[{"role": "user", "content": prompt}])
assert caught.value.status_code == 400
assert _GENERIC in str(caught.value), str(caught.value)
_only_call(wire)
async def test_async_openai_sdk_chat_completions_raises_bad_request_saying_prompt_too_long(
gateway: Gateway,
) -> None:
prompt: Final = f"overflow async sdk chat {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
async with openai.AsyncOpenAI(
base_url=f"{_proxy_url(gateway)}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(timeout=15, trust_env=False),
) as client:
with pytest.raises(openai.BadRequestError) as caught:
await client.chat.completions.create(model=model, messages=[{"role": "user", "content": prompt}])
assert caught.value.status_code == 400
assert _GENERIC in str(caught.value), str(caught.value)
_only_call(wire)
def test_openai_sdk_responses_raises_bad_request_saying_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow sdk responses {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
with pytest.raises(openai.BadRequestError) as caught:
_openai_client(gateway).responses.create(model=model, input=prompt)
assert caught.value.status_code == 400
assert _GENERIC in str(caught.value), str(caught.value)
_only_call(wire)
def test_anthropic_sdk_messages_raises_bad_request_saying_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow sdk messages {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
with pytest.raises(anthropic.BadRequestError) as caught:
_anthropic_client(gateway).messages.create(
model=model, max_tokens=32, messages=[{"role": "user", "content": prompt}]
)
assert caught.value.status_code == 400
assert _GENERIC in str(caught.value), str(caught.value)
_only_call(wire)
async def test_async_anthropic_sdk_messages_raises_bad_request_saying_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"overflow async sdk messages {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
async with anthropic.AsyncAnthropic(
base_url=_proxy_url(gateway),
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(timeout=15, trust_env=False),
) as client:
with pytest.raises(anthropic.BadRequestError) as caught:
await client.messages.create(model=model, max_tokens=32, messages=[{"role": "user", "content": prompt}])
assert caught.value.status_code == 400
assert _GENERIC in str(caught.value), str(caught.value)
_only_call(wire)
def test_anthropic_sdk_messages_stream_raises_invalid_request_error_saying_prompt_too_long(
gateway: Gateway,
) -> None:
prompt: Final = f"overflow sdk messages stream {uuid4().hex}"
reply: Final = _failed_stream(_OVERFLOW_STREAM_ERROR)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
with pytest.raises(anthropic.APIStatusError) as caught:
for _ in _anthropic_client(gateway).messages.create(
model=model, max_tokens=32, messages=[{"role": "user", "content": prompt}], stream=True
):
pass
body: Final = object_value(_JSON_VALUE.validate_python(caught.value.body))
error: Final = object_value(body["error"])
assert error["type"] == "invalid_request_error", body
assert _GENERIC in str(error["message"]), body
_only_call(wire)
def test_chat_completions_stream_overflow_in_response_failed_event_returns_400_prompt_too_long(
gateway: Gateway,
) -> None:
prompt: Final = f"failed event chat {uuid4().hex}"
reply: Final = _failed_stream(_OVERFLOW_STREAM_ERROR)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/chat/completions", _chat(model, prompt, stream=True))
assert response.status_code == 400, text
assert _GENERIC in text, text
_only_call(wire)
def test_messages_stream_overflow_in_response_failed_event_emits_invalid_request_error_event(
gateway: Gateway,
) -> None:
prompt: Final = f"failed event messages {uuid4().hex}"
reply: Final = _failed_stream(_OVERFLOW_STREAM_ERROR)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 200, text
assert response.headers["content-type"].startswith("text/event-stream"), dict(response.headers)
error: Final = _only_error(text)
assert error["type"] == "invalid_request_error", text
assert _GENERIC in str(error["message"]), text
assert not any(event["type"] == "content_block_delta" for event in _events(text)), text
_only_call(wire)
def test_responses_stream_overflow_in_response_failed_event_relays_failure_saying_prompt_too_long(
gateway: Gateway,
) -> None:
prompt: Final = f"failed event responses {uuid4().hex}"
reply: Final = _failed_stream(_OVERFLOW_STREAM_ERROR)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/responses", _responses(model, prompt, stream=True))
assert response.status_code == 200, text
assert _GENERIC in str(_only_failed(text)["message"]), text
_only_call(wire)
def test_chat_completions_non_overflow_400_keeps_the_upstream_message(gateway: Gateway) -> None:
prompt: Final = f"bad input chat {uuid4().hex}"
reply: Final = Reply(status=400, body=_BAD_INPUT_BODY)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 400, response.text
assert _INVALID_INPUT_MESSAGE in response.text, response.text
assert _TOO_LONG not in response.text, response.text
_only_call(wire)
def test_messages_stream_non_overflow_400_before_the_stream_returns_400_with_the_upstream_message(
gateway: Gateway,
) -> None:
prompt: Final = f"bad input messages stream {uuid4().hex}"
reply: Final = Reply(status=400, body=_BAD_INPUT_BODY)
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 400, text
assert _INVALID_INPUT_MESSAGE in text, text
assert _TOO_LONG not in text, text
_only_call(wire)
def test_messages_stream_non_overflow_response_failed_event_still_returns_500_api_error(gateway: Gateway) -> None:
prompt: Final = f"invalid prompt messages stream {uuid4().hex}"
reply: Final = _failed_stream({"code": "invalid_prompt", "message": _INVALID_PROMPT_MESSAGE})
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 500, text
error: Final = object_value(_JSON_OBJECT.validate_json(text)["error"])
assert error["type"] == "api_error", text
assert _INVALID_PROMPT_MESSAGE in str(error["message"]), text
assert _TOO_LONG not in text, text
_only_call(wire)
@pytest.mark.parametrize(
("status", "body", "expected"),
(
(401, b'{"error":{"message":"invalid bearer","type":"authentication_error","code":null}}', 401),
(429, b'{"error":{"message":"slow down","type":"rate_limit_error","code":"rate_limit_exceeded"}}', 429),
(500, b'{"error":{"message":"boom","type":"server_error","code":null}}', 503),
(
404,
b'{"error":{"message":"The model `x` does not exist","type":"invalid_request_error","code":"model_not_found"}}',
404,
),
),
ids=("401", "429", "500", "404"),
)
def test_chat_completions_other_upstream_statuses_keep_their_mapping(
gateway: Gateway, status: int, body: bytes, expected: int
) -> None:
prompt: Final = f"status {status} chat {uuid4().hex}"
with wire_server(_mantle_peer(prompt, Reply(status=status, body=body))) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == expected, response.text
assert _TOO_LONG not in response.text, response.text
_only_call(wire)
def test_messages_stream_legacy_token_count_envelope_returns_400_with_the_counts(gateway: Gateway) -> None:
prompt: Final = f"legacy messages stream {uuid4().hex}"
with (
wire_server(_mantle_peer(prompt, Reply(status=400, body=_LEGACY_BODY))) as wire,
gateway.scenario() as scenario,
):
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 400, text
error: Final = object_value(_JSON_OBJECT.validate_json(text)["error"])
assert error["type"] == "invalid_request_error", text
assert f"prompt is too long: {_LEGACY_PROMPT_TOKENS} tokens > {_LEGACY_MODEL_MAXIMUM} maximum" in str(
error["message"]
), text
_only_call(wire)
def test_chat_completions_code_only_overflow_envelope_returns_400_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"code only chat {uuid4().hex}"
reply: Final = Reply(status=400, body=b'{"error":{"code":"context_length_exceeded"}}')
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 400, response.text
assert _GENERIC in _error_message(response), response.text
_only_call(wire)
def test_chat_completions_text_plain_overflow_body_returns_400_prompt_too_long(gateway: Gateway) -> None:
prompt: Final = f"text plain chat {uuid4().hex}"
reply: Final = Reply(status=400, body=b"request rejected: context_length_exceeded", content_type="text/plain")
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 400, response.text
assert _GENERIC in _error_message(response), response.text
_only_call(wire)
def test_chat_completions_five_kilobyte_overflow_message_returns_400_and_keeps_the_proxy_alive(
gateway: Gateway,
) -> None:
prompt: Final = f"large envelope chat {uuid4().hex}"
reply: Final = Reply(status=400, body=_envelope("context_length_exceeded", "x" * 5120))
with wire_server(_mantle_peer(prompt, reply)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 400, response.text
assert _GENERIC in _error_message(response), response.text
_only_call(wire)
assert gateway.request("GET", "/health/liveliness").status_code == 200
def test_chat_completions_stream_response_failed_without_error_fails_and_keeps_the_proxy_alive(
gateway: Gateway,
) -> None:
prompt: Final = f"null error chat stream {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _failed_stream(None))) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response, text = _consumed(gateway, "/v1/chat/completions", _chat(model, prompt, stream=True))
assert response.status_code >= 400, text
assert _TOO_LONG not in text, text
_only_call(wire)
assert gateway.request("GET", "/health/liveliness").status_code == 200
def test_chat_completions_repeated_overflow_logs_one_failure_row_per_call(gateway: Gateway) -> None:
prompt: Final = f"repeated overflow chat {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _OVERFLOW)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
responses: Final = tuple(
gateway.request("POST", "/v1/chat/completions", _chat(model, prompt)) for _ in range(2)
)
call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses)
assert len(set(call_ids)) == 2, call_ids
assert [request.target for request in wire.drain()] == [_RESPONSES_PATH, _RESPONSES_PATH]
for response, call_id in zip(responses, call_ids, strict=True):
assert response.status_code == 400, response.text
assert _GENERIC in _failure_message(call_id)
def test_chat_completions_prompt_naming_the_error_code_still_succeeds(gateway: Gateway) -> None:
prompt: Final = f"my prompt mentions context_length_exceeded {uuid4().hex}"
with wire_server(_mantle_peer(prompt, _happy_reply(stream=False))) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request("POST", "/v1/chat/completions", _chat(model, prompt))
assert response.status_code == 200, response.text
assert _HAPPY_TEXT in response.text, response.text
_only_call(wire)
def test_messages_stream_openai_deployment_overflow_in_the_stream_emits_invalid_request_error_event(
gateway: Gateway,
) -> None:
prompt: Final = f"openai overflow messages stream {uuid4().hex}"
with wire_server(_openai_peer(prompt)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_OPENAI_MODEL, api_base=f"{wire.url}/v1", api_key=_OPENAI_API_KEY)
response, text = _consumed(gateway, "/v1/messages", _messages(model, prompt, stream=True))
assert response.status_code == 200, text
error: Final = _only_error(text)
assert error["type"] == "invalid_request_error", text
assert _OPENAI_OVERFLOW_MARK in str(error["message"]), text
_only_call(wire, _OPENAI_RESPONSES_PATH)
def test_chat_completions_stream_openai_deployment_overflow_returns_400(gateway: Gateway) -> None:
prompt: Final = f"openai overflow chat stream {uuid4().hex}"
with wire_server(_openai_peer(prompt)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_OPENAI_MODEL, api_base=f"{wire.url}/v1", api_key=_OPENAI_API_KEY)
response, text = _consumed(gateway, "/v1/chat/completions", _chat(model, prompt, stream=True))
assert response.status_code == 400, text
assert _OPENAI_OVERFLOW_MARK in text, text
_only_call(wire, _OPENAI_CHAT_PATH)
def test_messages_openai_deployment_overflow_returns_400_invalid_request(gateway: Gateway) -> None:
prompt: Final = f"openai overflow messages {uuid4().hex}"
with wire_server(_openai_peer(prompt)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_OPENAI_MODEL, api_base=f"{wire.url}/v1", api_key=_OPENAI_API_KEY)
response: Final = gateway.request("POST", "/v1/messages", _messages(model, prompt))
assert response.status_code == 400, response.text
error: Final = _error_object(response)
assert error["type"] == "invalid_request_error", response.text
assert _OPENAI_OVERFLOW_MARK in str(error["message"]), response.text
_only_call(wire, _OPENAI_RESPONSES_PATH)
_OVERFLOW_MARKER: Final = "chaos-overflow"
_HAPPY_MARKER: Final = "chaos-happy"
_Call = tuple[str, dict[str, JsonValue], bool, str]
_Outcome = tuple[str, bool, bool, int, str, str]
_Builder = Callable[[str, str, bool], dict[str, JsonValue]]
_Cell = tuple[tuple[str, _Builder], str, bool]
_BUILDERS: Final = (("/v1/chat/completions", _chat), ("/v1/messages", _messages), ("/v1/responses", _responses))
_STREAMED_MESSAGES_OVERFLOW: Final = ("/v1/messages", True, True)
def _logs_a_spend_row(path: str, overflow: bool, stream: bool) -> bool:
return (path, overflow, stream) != _STREAMED_MESSAGES_OVERFLOW
def _chaos_peer(request: Request) -> Reply:
body: Final = _JSON_OBJECT.validate_json(request.body)
stream: Final = body.get("stream") is True
if _OVERFLOW_MARKER in json.dumps(body["input"]):
return _failed_stream(_OVERFLOW_STREAM_ERROR) if stream else _OVERFLOW
return _happy_reply(stream)
def _burst_bodies(model: str, round_name: str) -> tuple[_Call, ...]:
def call(cell: _Cell) -> _Call:
(path, build), marker, stream = cell
tag: Final = f"chaos-{round_name}-{uuid4().hex}"
prompt: Final = f"{marker} {round_name} {path} stream={stream} {tag}"
return path, build(model, prompt, stream), marker == _OVERFLOW_MARKER, tag
return tuple(call(cell) for cell in product(_BUILDERS, (_HAPPY_MARKER, _OVERFLOW_MARKER), (False, True)))
def _tagged_rows(tag: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT status FROM "LiteLLM_SpendLogs" WHERE request_tags::jsonb @> %s::jsonb', (json.dumps([tag]),)
)
def _single_tagged_status(tag: str) -> str:
rows: Final = eventually(lambda: _tagged_rows(tag), lambda values: len(values) == 1, seconds=70)
return str(rows[0]["status"])
def _fire(gateway: Gateway, bodies: tuple[_Call, ...]) -> tuple[_Outcome, ...]:
def one(item: _Call) -> _Outcome:
path, body, overflow, tag = item
with gateway.client.stream(
"POST", path, json=body, headers={"Authorization": f"Bearer {gateway.key}", "x-litellm-tags": tag}
) as response:
text: Final = b"".join(response.iter_bytes()).decode()
return path, overflow, body["stream"] is True, response.status_code, tag, text
with ThreadPoolExecutor(max_workers=len(bodies)) as pool:
return tuple(pool.map(one, bodies))
def _assert_served(outcomes: tuple[_Outcome, ...]) -> None:
for path, overflow, stream, status, tag, text in outcomes:
if not overflow:
assert status == 200, (path, text)
assert _HAPPY_TEXT in text, (path, text)
assert _single_tagged_status(tag) == "success", (path, tag)
continue
assert _GENERIC in text, (path, status, text)
assert status in (200, 400), (path, status, text)
if _logs_a_spend_row(path, overflow, stream):
assert _single_tagged_status(tag) == "failure", (path, tag)
def test_messages_stream_overflow_logs_one_failure_row(gateway: Gateway) -> None:
pytest.skip("BUG: a streamed /v1/messages context overflow writes no LiteLLM_SpendLogs row (LIT-9132)")
with wire_server(_chaos_peer) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
tag: Final = f"chaos-single-{uuid4().hex}"
body: Final = _messages(model, f"{_OVERFLOW_MARKER} single {tag}", True)
((_, _, _, status, _, text),) = _fire(gateway, (("/v1/messages", body, True, tag),))
assert status == 200, text
assert _GENERIC in text, text
assert _single_tagged_status(tag) == "failure", tag
def test_chaos_mantle_peer_outage_mid_burst_logs_every_call_once_and_keeps_the_proxy_alive(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
with wire_server(_chaos_peer) as wire:
port: Final = int(wire.url.rsplit(":", 1)[1])
model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_API_KEY)
before: Final = _fire(gateway, _burst_bodies(model, "before"))
assert len(wire.drain()) == len(before)
during: Final = _fire(gateway, _burst_bodies(model, "during"))
assert gateway.request("GET", "/health/liveliness").status_code == 200
with wire_server(_chaos_peer, port=port) as revived:
after: Final = _fire(gateway, _burst_bodies(model, "after"))
assert len(revived.drain()) == len(after)
_assert_served(before)
_assert_served(after)
for path, _, _, status, tag, text in during:
assert status >= 500 or _error_events(text) or '"response.failed"' in text, (path, status, text)
assert _TOO_LONG not in text, (path, text)
assert _single_tagged_status(tag) == "failure", (path, tag)

View file

@ -1181,6 +1181,47 @@ def test_bedrock_mantle_context_overflow_maps_to_context_window_exceeded():
assert "prompt is too long: 1055489 tokens > 1050000 maximum" in excinfo.value.message
def test_bedrock_mantle_openai_envelope_context_overflow_maps_to_context_window_exceeded():
from litellm.llms.base_llm.chat.transformation import BaseLLMException
original_exception = BaseLLMException(
status_code=400,
message=(
'{"error":{"code":"context_length_exceeded",'
'"message":"Your input exceeds the context window of this model. '
'Please adjust your input and try again.",'
'"param":"input","type":"invalid_request_error"}}'
),
)
with pytest.raises(litellm.ContextWindowExceededError) as excinfo:
exception_type(
model="openai.gpt-5.6-sol",
original_exception=original_exception,
custom_llm_provider="bedrock_mantle",
)
assert excinfo.value.status_code == 400
assert "prompt is too long" in excinfo.value.message
def test_bedrock_mantle_streamed_context_overflow_event_maps_to_context_window_exceeded():
from litellm.responses.streaming_iterator import _map_stream_error_to_exception
mapped_exception = _map_stream_error_to_exception(
{
"code": "context_length_exceeded",
"message": "Your input exceeds the context window of this model. Please adjust your input and try again.",
},
model="openai.gpt-5.6-luna",
custom_llm_provider="bedrock_mantle",
)
assert isinstance(mapped_exception, litellm.ContextWindowExceededError)
assert mapped_exception.status_code == 400
assert "prompt is too long" in mapped_exception.message
def test_branchless_provider_transport_error_maps_to_api_connection_error():
from litellm.llms.base_llm.chat.transformation import BaseLLMException

View file

@ -22,6 +22,8 @@ from unittest.mock import MagicMock
import pytest
import litellm
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.exceptions import MidStreamFallbackError
@ -140,3 +142,31 @@ def test_error_event_preserves_midstream_fallback_error():
assert name == "error"
assert payload["error"]["type"] == "api_error"
assert "internalServerException" in payload["error"]["message"]
def test_error_event_keeps_status_of_mapped_litellm_exception():
exc = litellm.ContextWindowExceededError(
message="prompt is too long: your prompt exceeds the model's context window",
model="openai.gpt-5.6-luna",
llm_provider="bedrock_mantle",
)
name, payload = _parse_sse(_mid_stream_error_sse_event(exc))
assert name == "error"
assert payload["error"]["type"] == "invalid_request_error"
assert "prompt is too long" in payload["error"]["message"]
@pytest.mark.parametrize(
"exc",
[
litellm.BadRequestError(
message="temperature must be in the range [0.0, 2.0]", model="m", llm_provider="gemini"
),
litellm.AuthenticationError(message="API key not valid", llm_provider="gemini", model="m"),
litellm.NotFoundError(message="model is not found", model="m", llm_provider="gemini"),
],
)
def test_error_event_reports_other_provider_4xx_as_retriable_500(exc):
name, payload = _parse_sse(_mid_stream_error_sse_event(exc))
assert name == "error"
assert payload["error"]["type"] == "api_error"