mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 750ba1e35a into 902736bfe7
This commit is contained in:
commit
23015e8d65
6 changed files with 1400 additions and 11 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue