diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 0fdfb301291..5a005d8059b 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -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( diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 4ef6c305cf5..d3847252835 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -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" diff --git a/tests/integration/providers/test_bedrock_mantle_context_overflow_fallbacks_wire.py b/tests/integration/providers/test_bedrock_mantle_context_overflow_fallbacks_wire.py new file mode 100644 index 00000000000..2b0c7f89c2c --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_context_overflow_fallbacks_wire.py @@ -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 diff --git a/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py new file mode 100644 index 00000000000..ad04535c309 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py @@ -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) diff --git a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py index 9de768ea47b..b347707a416 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -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 diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py index 4798d522182..259b57b9270 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_mid_stream_error.py @@ -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"