From b4f9e14a44789b408c3ceab56d01671f33a7b3d0 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 28 Jul 2026 14:14:53 +0000 Subject: [PATCH 01/12] feat(dd_span_tagger): emit litellm_user_email span tag for JWT-authenticated requests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/dd_span_tagger.py | 8 +++++- .../proxy/test_common_request_processing.py | 27 ++++++++++++++++++- 2 files changed, 33 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/dd_span_tagger.py b/litellm/proxy/dd_span_tagger.py index 08b7d928d0e..1fda0c638d4 100644 --- a/litellm/proxy/dd_span_tagger.py +++ b/litellm/proxy/dd_span_tagger.py @@ -38,11 +38,15 @@ class DDSpanTagger: - ``litellm.key_alias`` — human-readable alias for the API key - ``litellm.key_hash`` — hashed API key (safe to log; never the raw secret) - ``litellm.requested_model``— model name as sent by the client + - ``litellm_user_email`` — email of the authenticated user Use cases: - Trace all requests from a specific user/key: filter by ``litellm.key_alias`` or ``litellm.key_hash``. - Trace all requests for a specific model: filter by ``litellm.requested_model``. + - Trace all requests from a specific person under JWT auth, where there is no key_alias: + filter by ``@litellm_user_email:"user@example.com"``. The value comes from the JWT claim + mapped by ``litellm_jwtauth.user_email_jwt_field``, or from the user row for virtual keys. Note: key_alias / key_hash are not available for unauthenticated (e.g. 401) requests. """ @@ -53,8 +57,10 @@ class DDSpanTagger: set_active_span_tag("litellm.key_hash", str(user_api_key_dict.token)) if requested_model: set_active_span_tag("litellm.requested_model", str(requested_model)) + if user_api_key_dict.user_email: + set_active_span_tag("litellm_user_email", str(user_api_key_dict.user_email)) except Exception: verbose_proxy_logger.debug( - "Failed to tag active ddtrace span with key/model tags", + "Failed to tag active ddtrace span with key/model/user tags", exc_info=True, ) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 58f81cdad35..287e6f2f80c 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2318,12 +2318,13 @@ class TestStreamingOverheadHeader: class TestDDSpanTaggerTagRequest: """Tests for DDSpanTagger.tag_request - key/model DD span tagging.""" - def _make_user_api_key_dict(self, key_alias=None, token=None): + def _make_user_api_key_dict(self, key_alias=None, token=None, user_email=None): from litellm.proxy._types import UserAPIKeyAuth d = UserAPIKeyAuth() d.key_alias = key_alias d.token = token + d.user_email = user_email return d def test_tags_key_alias_and_model(self): @@ -2364,6 +2365,30 @@ class TestDDSpanTaggerTagRequest: mock_set_tag.assert_called_once_with("litellm.requested_model", "claude-3-5-sonnet") + def test_tags_user_email(self): + """user_email is tagged so JWT-authenticated requests are traceable per person.""" + user_key = self._make_user_api_key_dict(user_email="user@example.com") + + with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag: + DDSpanTagger.tag_request( + user_api_key_dict=user_key, + requested_model=None, + ) + + mock_set_tag.assert_called_once_with("litellm_user_email", "user@example.com") + + def test_no_user_email_tag_when_absent(self): + """No user email tag when the authenticated identity has no email.""" + user_key = self._make_user_api_key_dict(key_alias="my-prod-key", user_email=None) + + with patch("litellm.proxy.dd_span_tagger.set_active_span_tag") as mock_set_tag: + DDSpanTagger.tag_request( + user_api_key_dict=user_key, + requested_model="gpt-4o", + ) + + assert all(call.args[0] != "litellm_user_email" for call in mock_set_tag.call_args_list) + class TestHasAttributeErrorInChain: """Tests for _has_attribute_error_in_chain helper.""" From 819faa01c336d825619b3592ed1c1d763a4156d4 Mon Sep 17 00:00:00 2001 From: milan Date: Tue, 28 Jul 2026 14:40:04 +0000 Subject: [PATCH 02/12] refactor(dd_span_tagger): use dotted litellm.user_email tag for consistency Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/dd_span_tagger.py | 6 +++--- tests/test_litellm/proxy/test_common_request_processing.py | 4 ++-- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/dd_span_tagger.py b/litellm/proxy/dd_span_tagger.py index 1fda0c638d4..cafae9bb2a3 100644 --- a/litellm/proxy/dd_span_tagger.py +++ b/litellm/proxy/dd_span_tagger.py @@ -38,14 +38,14 @@ class DDSpanTagger: - ``litellm.key_alias`` — human-readable alias for the API key - ``litellm.key_hash`` — hashed API key (safe to log; never the raw secret) - ``litellm.requested_model``— model name as sent by the client - - ``litellm_user_email`` — email of the authenticated user + - ``litellm.user_email`` — email of the authenticated user Use cases: - Trace all requests from a specific user/key: filter by ``litellm.key_alias`` or ``litellm.key_hash``. - Trace all requests for a specific model: filter by ``litellm.requested_model``. - Trace all requests from a specific person under JWT auth, where there is no key_alias: - filter by ``@litellm_user_email:"user@example.com"``. The value comes from the JWT claim + filter by ``@litellm.user_email:"user@example.com"``. The value comes from the JWT claim mapped by ``litellm_jwtauth.user_email_jwt_field``, or from the user row for virtual keys. Note: key_alias / key_hash are not available for unauthenticated (e.g. 401) requests. @@ -58,7 +58,7 @@ class DDSpanTagger: if requested_model: set_active_span_tag("litellm.requested_model", str(requested_model)) if user_api_key_dict.user_email: - set_active_span_tag("litellm_user_email", str(user_api_key_dict.user_email)) + set_active_span_tag("litellm.user_email", str(user_api_key_dict.user_email)) except Exception: verbose_proxy_logger.debug( "Failed to tag active ddtrace span with key/model/user tags", diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 287e6f2f80c..312c3d2a884 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -2375,7 +2375,7 @@ class TestDDSpanTaggerTagRequest: requested_model=None, ) - mock_set_tag.assert_called_once_with("litellm_user_email", "user@example.com") + mock_set_tag.assert_called_once_with("litellm.user_email", "user@example.com") def test_no_user_email_tag_when_absent(self): """No user email tag when the authenticated identity has no email.""" @@ -2387,7 +2387,7 @@ class TestDDSpanTaggerTagRequest: requested_model="gpt-4o", ) - assert all(call.args[0] != "litellm_user_email" for call in mock_set_tag.call_args_list) + assert all(call.args[0] != "litellm.user_email" for call in mock_set_tag.call_args_list) class TestHasAttributeErrorInChain: From b224b15b9d24c15590a81c3c4bb5399d40b7d9e0 Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 02:41:55 +0000 Subject: [PATCH 03/12] feat(model_hub): surface model_info.description in model group info and Model Hub UI Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/router.py | 6 +++++ litellm/types/router.py | 1 + tests/test_litellm/test_router.py | 25 +++++++++++++++++++ .../src/components/AIHub/ModelHubTable.tsx | 6 +++++ .../components/AIHub/ModelHubTableColumns.tsx | 1 + .../components/PublicModelHubTableColumns.tsx | 1 + .../src/components/public_model_hub.tsx | 6 +++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 ++ 8 files changed, 48 insertions(+) diff --git a/litellm/router.py b/litellm/router.py index e9eeab53934..90968800460 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9334,6 +9334,9 @@ class Router: model_litellm_params = model.get("litellm_params", {}) model_info_dict = model.get("model_info", {}) + _raw_description = model_info_dict.get("description") + _deployment_description: str | None = _raw_description if isinstance(_raw_description, str) else None + # get model tpm _deployment_tpm: int | None = None if _deployment_tpm is None: @@ -9415,6 +9418,7 @@ class Router: **{ "model_group": user_facing_model_group_name, "providers": [llm_provider], + "description": _deployment_description, **model_info, } ) @@ -9428,6 +9432,8 @@ class Router: # supports_function_calling == True if llm_provider not in model_group_info.providers: model_group_info.providers.append(llm_provider) + if model_group_info.description is None and _deployment_description is not None: + model_group_info.description = _deployment_description if ( model_info.get("max_input_tokens", None) is not None and model_info["max_input_tokens"] is not None diff --git a/litellm/types/router.py b/litellm/types/router.py index 99a4603ae49..b33e387c508 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -637,6 +637,7 @@ class ModelGroupInfo(BaseModel): supports_function_calling: bool = Field(default=False) supported_openai_params: list[str] | None = Field(default=[]) configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS = None + description: str | None = None def __init__(self, **data) -> None: for field_name, field_type in get_type_hints(self.__class__).items(): diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 9894fcef163..765e479e378 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -1626,6 +1626,31 @@ def test_model_group_info_cost_none_when_db_model_info_has_no_cost(): assert result.output_cost_per_token is None +def test_model_group_info_description_from_model_info(): + """ + model_info.description set in the config should surface on ModelGroupInfo, + including when only a later deployment in the group carries it. + """ + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake"}, + "model_info": {}, + }, + { + "model_name": "gpt-4", + "litellm_params": {"model": "gpt-4", "api_key": "fake2"}, + "model_info": {"description": "State-of-the-art language model."}, + }, + ] + ) + + result = router.get_model_group_info("gpt-4") + assert result is not None + assert result.description == "State-of-the-art language model." + + @pytest.mark.parametrize( "value,expected", [ diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index 299104b271b..5fadbe3300d 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -647,6 +647,12 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, ))} + {selectedModel.description && ( +
+

Description:

+

{selectedModel.description}

+
+ )} diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx index 9f74771f3b1..b025bcf62aa 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTableColumns.tsx @@ -31,6 +31,7 @@ export interface ModelHubData { supports_function_calling: boolean; supported_openai_params?: string[]; is_public_model_group: boolean; + description?: string; [key: string]: any; } diff --git a/ui/litellm-dashboard/src/components/PublicModelHubTableColumns.tsx b/ui/litellm-dashboard/src/components/PublicModelHubTableColumns.tsx index ab0ed976149..f3673df8025 100644 --- a/ui/litellm-dashboard/src/components/PublicModelHubTableColumns.tsx +++ b/ui/litellm-dashboard/src/components/PublicModelHubTableColumns.tsx @@ -24,6 +24,7 @@ export interface ModelGroupInfo { health_status?: string; health_response_time?: number; health_checked_at?: string; + description?: string; [key: string]: any; } diff --git a/ui/litellm-dashboard/src/components/public_model_hub.tsx b/ui/litellm-dashboard/src/components/public_model_hub.tsx index 171b7992325..cd10068b9e6 100644 --- a/ui/litellm-dashboard/src/components/public_model_hub.tsx +++ b/ui/litellm-dashboard/src/components/public_model_hub.tsx @@ -918,6 +918,12 @@ const PublicModelHub: React.FC = ({ accessToken, isEmbedded })} + {selectedModel.description && ( +
+

Description:

+

{selectedModel.description}

+
+ )} {/* Wildcard Routing Note */} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index fff48e14ecf..39b3255efb3 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -29481,6 +29481,8 @@ export interface components { ModelGroupInfoProxy: { /** Configurable Clientside Auth Params */ configurable_clientside_auth_params?: (string | components["schemas"]["ConfigurableClientsideParamsCustomAuth-Output"])[] | null; + /** Description */ + description?: string | null; /** Health Checked At */ health_checked_at?: string | null; /** Health Response Time */ From 342b93fe0a92f32b8c2fb4ccf61441cb0a2cf3b1 Mon Sep 17 00:00:00 2001 From: milan Date: Fri, 21 Aug 2026 03:27:35 +0000 Subject: [PATCH 04/12] refactor(router): aggregate model group description without in-place mutation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/router.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 90968800460..d65c13d57b2 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9313,6 +9313,10 @@ class Router: model_list: Final = self.get_model_list(model_name=model_group) if model_list is None: return None + model_group_description: Final[str | None] = next( + (d for d in (m.get("model_info", {}).get("description") for m in model_list) if isinstance(d, str)), + None, + ) for model in model_list: is_match = False if ( @@ -9334,9 +9338,6 @@ class Router: model_litellm_params = model.get("litellm_params", {}) model_info_dict = model.get("model_info", {}) - _raw_description = model_info_dict.get("description") - _deployment_description: str | None = _raw_description if isinstance(_raw_description, str) else None - # get model tpm _deployment_tpm: int | None = None if _deployment_tpm is None: @@ -9418,7 +9419,7 @@ class Router: **{ "model_group": user_facing_model_group_name, "providers": [llm_provider], - "description": _deployment_description, + "description": model_group_description, **model_info, } ) @@ -9432,8 +9433,6 @@ class Router: # supports_function_calling == True if llm_provider not in model_group_info.providers: model_group_info.providers.append(llm_provider) - if model_group_info.description is None and _deployment_description is not None: - model_group_info.description = _deployment_description if ( model_info.get("max_input_tokens", None) is not None and model_info["max_input_tokens"] is not None From 6d2c4899b0f9fa36c59c793657234f75b6a83704 Mon Sep 17 00:00:00 2001 From: yassin Date: Sun, 13 Sep 2026 09:44:22 +0000 Subject: [PATCH 05/12] fix(health): skip background health check DB writes when the latest-row read fails Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/db/health_check_latest.py | 8 ++++-- .../health_endpoints/_health_endpoints.py | 7 +++-- .../proxy/db/test_health_check_latest.py | 10 +++++++ .../proxy/test_health_check_functions.py | 27 +++++++++++++++---- 4 files changed, 43 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/db/health_check_latest.py b/litellm/proxy/db/health_check_latest.py index 35bc838379c..21438f095bb 100644 --- a/litellm/proxy/db/health_check_latest.py +++ b/litellm/proxy/db/health_check_latest.py @@ -74,10 +74,14 @@ class LatestHealthCheckRow(BaseModel): _ROWS_ADAPTER: Final = TypeAdapter(tuple[LatestHealthCheckRow, ...]) +async def query_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: + rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) + return _ROWS_ADAPTER.validate_python(rows) + + async def fetch_latest_health_checks(prisma_client: PrismaClient) -> tuple[LatestHealthCheckRow, ...]: try: - rows: Final = await prisma_client.db.query_raw(LATEST_HEALTH_CHECKS_SQL) - return _ROWS_ADAPTER.validate_python(rows) + return await query_latest_health_checks(prisma_client) except Exception as query_err: # noqa: BLE001 # health decorates other reads; a driver error must not fail them verbose_proxy_logger.error("Error getting all latest health checks: %s", query_err) return () diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index db6ec754c6e..64fd59bbe44 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -45,7 +45,10 @@ from litellm.proxy.auth.auth_utils import ( from litellm.proxy.auth.model_checks import get_key_models from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler -from litellm.proxy.db.health_check_latest import LatestHealthCheckRow +from litellm.proxy.db.health_check_latest import ( + LatestHealthCheckRow, + query_latest_health_checks, +) from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers from litellm.proxy.health_check import ( ADMIN_ONLY_HEALTH_DISPLAY_PARAMS, @@ -876,7 +879,7 @@ async def _save_background_health_checks_to_db( ) # Step 3: Get latest health checks for all models in one query to compare status - latest_checks: Final = await prisma_client.get_all_latest_health_checks() + latest_checks: Final = await query_latest_health_checks(prisma_client) latest_checks_map: Final = {} for check in latest_checks: # Use model_id as primary key, fallback to model_name diff --git a/tests/test_litellm/proxy/db/test_health_check_latest.py b/tests/test_litellm/proxy/db/test_health_check_latest.py index 529e4d8f2e9..6322891ae9e 100644 --- a/tests/test_litellm/proxy/db/test_health_check_latest.py +++ b/tests/test_litellm/proxy/db/test_health_check_latest.py @@ -8,6 +8,7 @@ from litellm.proxy.db.health_check_latest import ( LATEST_HEALTH_CHECKS_SQL, fetch_latest_health_checks, fetch_latest_health_checks_for_models, + query_latest_health_checks, ) @@ -83,6 +84,15 @@ async def test_fetch_all_degrades_to_no_rows_when_the_query_fails(): assert await fetch_latest_health_checks(prisma) == () +@pytest.mark.asyncio +async def test_query_all_raises_when_the_query_fails_instead_of_reading_as_an_empty_table(): + """The background save decides what to write from this read; a failure has to be told apart from no rows.""" + prisma = _prisma([]) + prisma.db.query_raw.side_effect = RuntimeError("db down") + with pytest.raises(RuntimeError, match="db down"): + await query_latest_health_checks(prisma) + + @pytest.mark.asyncio async def test_fetch_all_degrades_to_no_rows_for_a_malformed_row(): assert await fetch_latest_health_checks(_prisma([{"unexpected": "shape"}])) == () diff --git a/tests/test_litellm/proxy/test_health_check_functions.py b/tests/test_litellm/proxy/test_health_check_functions.py index c0c853ae2c5..7d6c5d3cebe 100644 --- a/tests/test_litellm/proxy/test_health_check_functions.py +++ b/tests/test_litellm/proxy/test_health_check_functions.py @@ -374,7 +374,7 @@ async def test_save_background_health_checks_to_db(): """Test the main background health check save function""" mock_prisma = MagicMock() mock_prisma.save_health_check_result = AsyncMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) model_list = [ { @@ -398,9 +398,9 @@ async def test_save_background_health_checks_to_db(): "background_health_check", ) - # Should call get_all_latest_health_checks and save_health_check_result, and report completion + # Should read the latest rows and save_health_check_result, and report completion assert persisted is True - mock_prisma.get_all_latest_health_checks.assert_called_once() + mock_prisma.db.query_raw.assert_awaited_once() mock_prisma.save_health_check_result.assert_called_once() call_kwargs = mock_prisma.save_health_check_result.call_args[1] @@ -493,7 +493,7 @@ def _one_model_setup(): @pytest.mark.asyncio async def test_save_background_health_checks_to_db_returns_false_when_a_write_fails(): mock_prisma = MagicMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.save_health_check_result = AsyncMock(return_value=None) model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() @@ -504,6 +504,23 @@ async def test_save_background_health_checks_to_db_returns_false_when_a_write_fa assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 1) +@pytest.mark.asyncio +async def test_save_background_health_checks_to_db_writes_nothing_when_the_latest_row_read_fails(mock_prisma): + """ + A failed dedup read must not read as an empty table. Treated that way, every model was written on every + cycle by every pod while the read kept failing, which is what filled the table in production. + """ + mock_prisma.db.query_raw = AsyncMock(side_effect=RuntimeError("db down")) + mock_prisma.save_health_check_result = AsyncMock(return_value={"id": "row"}) + model_list, healthy_endpoints, unhealthy_endpoints = _one_model_setup() + + persisted = await _save_background_health_checks_to_db( + mock_prisma, model_list, healthy_endpoints, unhealthy_endpoints, 1234567890.0, "background_health_check" + ) + + assert (persisted, mock_prisma.save_health_check_result.await_count) == (False, 0) + + @pytest.mark.asyncio async def test_save_background_health_checks_to_db_no_prisma(): """Test graceful handling when no prisma client""" @@ -515,7 +532,7 @@ async def test_save_background_health_checks_to_db_no_prisma(): async def test_save_background_health_checks_to_db_exception_handling(): """Test exception handling in background health check save""" mock_prisma = MagicMock() - mock_prisma.get_all_latest_health_checks = AsyncMock(side_effect=Exception("DB Error")) + mock_prisma.db.query_raw = AsyncMock(side_effect=Exception("DB Error")) model_list = [ { From 83049e2103bb0870f5521fb7ca4242e1efda28b9 Mon Sep 17 00:00:00 2001 From: mayank-affirm <78003505+mayank-affirm@users.noreply.github.com> Date: Thu, 17 Sep 2026 22:20:47 -0500 Subject: [PATCH 06/12] fix(bedrock_mantle): map the OpenAI-shaped context overflow envelope too Mantle reports context overflow in two envelopes. #37862 covered the structured `validation_error` one carrying token counts ("prompt tokens (N) exceed model maximum (M)"). Mantle also returns the OpenAI-shaped body: {"error":{"code":"context_length_exceeded","message":"Your input exceeds the context window of this model. ...","param":"input", "type":"invalid_request_error"}} That passes the existing `invalid_request_error` gate but matches neither the token-count pattern nor `is_error_str_context_window_exceeded`'s substring list, so it fell through to a generic BadRequestError. Clients such as Claude Code only reactive-compact on a recognised overflow phrase, so an overflowing request surfaced as an ordinary error with no recovery. Return the same normalised "prompt is too long: ..." message for this envelope when it carries no counts, keeping the numeric wording when the provider supplies them. --- .../exception_mapping_utils.py | 21 ++++++++---- .../test_exception_mapping_utils.py | 32 +++++++++++++++++++ 2 files changed, 46 insertions(+), 7 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 82708d412c9..19e33500e83 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -826,21 +826,28 @@ 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_OPENAI_PATTERN: Final = re.compile(r"exceeds the context window") +_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 in two envelopes, neither of which is the + plain-text wording Bedrock itself uses, so it needs its own detection and a + message clients such as Claude Code recognize as context 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 _BEDROCK_MANTLE_CONTEXT_WINDOW_OPENAI_PATTERN.search(error_str): + return _BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE + return None def _map_bedrock_exception( diff --git a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py index 42d3df76902..d0ebe245a9f 100644 --- a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py @@ -1209,6 +1209,38 @@ 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(): + """Mantle's OpenAI-style overflow envelope is a second, separate wording. + + Mantle returns context overflow either as a structured ``validation_error`` + carrying token counts, or as this OpenAI-shaped + ``context_length_exceeded`` body. Only the first was matched, so an overflow + in the second shape reached callers as a generic ``BadRequestError`` that + Claude Code cannot recognise, and its reactive compaction never retried. + """ + 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_branchless_provider_transport_error_maps_to_api_connection_error(): from litellm.llms.base_llm.chat.transformation import BaseLLMException From 51097e55ab4ed31eca1f7f129992b3640e5a13ab Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 19:33:40 -0700 Subject: [PATCH 07/12] fix(bedrock_mantle): map the streamed context overflow and keep its 400 on /v1/messages streams --- .../exception_mapping_utils.py | 13 ++++------ .../adapters/streaming_iterator.py | 4 +-- .../test_exception_mapping_utils.py | 25 +++++++++++++------ ...est_streaming_iterator_mid_stream_error.py | 14 +++++++++++ 4 files changed, 38 insertions(+), 18 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 554e90c2b3e..5a005d8059b 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -852,7 +852,6 @@ 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_OPENAI_PATTERN: Final = re.compile(r"exceeds the context window") _BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE: Final = ( "prompt is too long: your prompt exceeds the model's context window" ) @@ -860,18 +859,16 @@ _BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE: Final = ( def _get_bedrock_mantle_context_window_message(error_str: str) -> str | None: """ - Mantle reports context overflow in two envelopes, neither of which is the - plain-text wording Bedrock itself uses, so it needs its own detection and a - message clients such as Claude Code 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 not None: prompt_tokens, max_tokens = match.groups() return f"prompt is too long: {prompt_tokens} tokens > {max_tokens} maximum" - if _BEDROCK_MANTLE_CONTEXT_WINDOW_OPENAI_PATTERN.search(error_str): + if "context_length_exceeded" in error_str: return _BEDROCK_MANTLE_CONTEXT_WINDOW_GENERIC_MESSAGE return None diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 12eee663ca5..03fcdfcf1b5 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -14,11 +14,11 @@ from typing import ( get_args, ) +from openai import APIStatusError from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm._uuid import uuid -from litellm.exceptions import MidStreamFallbackError from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.anthropic import ( AppliedEdit, @@ -61,7 +61,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, APIStatusError)): return exc.status_code, exc.message return 500, str(exc) or "Upstream stream ended before completion" 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 55eff02179c..b347707a416 100644 --- a/tests/unit/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/unit/litellm_core_utils/test_exception_mapping_utils.py @@ -1182,14 +1182,6 @@ def test_bedrock_mantle_context_overflow_maps_to_context_window_exceeded(): def test_bedrock_mantle_openai_envelope_context_overflow_maps_to_context_window_exceeded(): - """Mantle's OpenAI-style overflow envelope is a second, separate wording. - - Mantle returns context overflow either as a structured ``validation_error`` - carrying token counts, or as this OpenAI-shaped - ``context_length_exceeded`` body. Only the first was matched, so an overflow - in the second shape reached callers as a generic ``BadRequestError`` that - Claude Code cannot recognise, and its reactive compaction never retried. - """ from litellm.llms.base_llm.chat.transformation import BaseLLMException original_exception = BaseLLMException( @@ -1213,6 +1205,23 @@ def test_bedrock_mantle_openai_envelope_context_overflow_maps_to_context_window_ 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..2a93a01f70e 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,15 @@ 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"] From b51bc2ed6d22b69c02c924fe3348c9ef7320cb10 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 29 Sep 2026 09:40:35 -0700 Subject: [PATCH 08/12] fix(anthropic_adapter): keep only the context overflow's 400 on streamed /v1/messages so other provider 4xx still fall back --- .../pass_through/adapters/streaming_iterator.py | 4 ++-- .../test_streaming_iterator_mid_stream_error.py | 14 ++++++++++++++ 2 files changed, 16 insertions(+), 2 deletions(-) diff --git a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py index 03fcdfcf1b5..0e09940125c 100644 --- a/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/pass_through/adapters/streaming_iterator.py @@ -14,11 +14,11 @@ from typing import ( get_args, ) -from openai import APIStatusError from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm._uuid import uuid +from litellm.exceptions import ContextWindowExceededError, MidStreamFallbackError from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.anthropic import ( AppliedEdit, @@ -61,7 +61,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, APIStatusError)): + 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/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 2a93a01f70e..7213493aef3 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 @@ -154,3 +154,17 @@ def test_error_event_keeps_status_of_mapped_litellm_exception(): 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" From 6f994f0360e5bf31e2603cbc2f7cca99a59d0fa3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 20:54:03 -0700 Subject: [PATCH 09/12] test(anthropic_adapter): wrap a parametrize line over the 120 char limit --- .../adapters/test_streaming_iterator_mid_stream_error.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) 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 7213493aef3..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 @@ -159,7 +159,9 @@ def test_error_event_keeps_status_of_mapped_litellm_exception(): @pytest.mark.parametrize( "exc", [ - litellm.BadRequestError(message="temperature must be in the range [0.0, 2.0]", model="m", llm_provider="gemini"), + 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"), ], From e158ec40ce2e5f32b9b584ee80d9997490dac577 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:07:57 -0700 Subject: [PATCH 10/12] test(integration): audit Bedrock Mantle context overflow mapping across routes, clients, and fallbacks --- ..._mantle_context_overflow_fallbacks_wire.py | 548 +++++++++++++ ...st_bedrock_mantle_context_overflow_wire.py | 749 ++++++++++++++++++ 2 files changed, 1297 insertions(+) create mode 100644 tests/integration/providers/test_bedrock_mantle_context_overflow_fallbacks_wire.py create mode 100644 tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py 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..6fa45c3b279 --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py @@ -0,0 +1,749 @@ +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] +_BUILDERS: Final = (("/v1/chat/completions", _chat), ("/v1/messages", _messages), ("/v1/responses", _responses)) +_UNLOGGED_OVERFLOW_CELL: Final = ("/v1/messages", True) + + +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: tuple[tuple[str, Callable[[str, str, bool], dict[str, JsonValue]]], str, bool]) -> _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 (path, stream) == _UNLOGGED_OVERFLOW_CELL: + continue + assert _single_tagged_status(tag) == "failure", (path, 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) From e28ad44dc706985b4f837a68381de918263b0ad0 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:27:40 -0700 Subject: [PATCH 11/12] test(integration): move the streamed /v1/messages overflow spend row into a BUG-skipped test --- ...st_bedrock_mantle_context_overflow_wire.py | 28 +++++++++++++++---- 1 file changed, 23 insertions(+), 5 deletions(-) diff --git a/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py index 6fa45c3b279..fab89ec7c28 100644 --- a/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py @@ -669,8 +669,15 @@ _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)) -_UNLOGGED_OVERFLOW_CELL: Final = ("/v1/messages", True) +_STREAMED_MESSAGES_OVERFLOW: Final = ("/v1/messages", _OVERFLOW_MARKER, True) + + +def _logs_a_spend_row(cell: _Cell) -> bool: + (path, _), marker, stream = cell + return (path, marker, stream) != _STREAMED_MESSAGES_OVERFLOW def _chaos_peer(request: Request) -> Reply: @@ -682,13 +689,14 @@ def _chaos_peer(request: Request) -> Reply: def _burst_bodies(model: str, round_name: str) -> tuple[_Call, ...]: - def call(cell: tuple[tuple[str, Callable[[str, str, bool], dict[str, JsonValue]]], str, bool]) -> _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))) + cells: Final = product(_BUILDERS, (_HAPPY_MARKER, _OVERFLOW_MARKER), (False, True)) + return tuple(call(cell) for cell in cells if _logs_a_spend_row(cell)) def _tagged_rows(tag: str) -> list[dict[str, JsonValue]]: @@ -724,11 +732,21 @@ def _assert_served(outcomes: tuple[_Outcome, ...]) -> None: continue assert _GENERIC in text, (path, status, text) assert status in (200, 400), (path, status, text) - if (path, stream) == _UNLOGGED_OVERFLOW_CELL: - continue 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: From 750ba1e35ad6b3a3d4bb6691a7cf8e767cf7a818 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 1 Oct 2026 17:46:48 -0700 Subject: [PATCH 12/12] test(integration): keep every chaos cell in the burst and skip only the missing spend row --- .../test_bedrock_mantle_context_overflow_wire.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py index fab89ec7c28..ad04535c309 100644 --- a/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_context_overflow_wire.py @@ -672,12 +672,11 @@ _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", _OVERFLOW_MARKER, True) +_STREAMED_MESSAGES_OVERFLOW: Final = ("/v1/messages", True, True) -def _logs_a_spend_row(cell: _Cell) -> bool: - (path, _), marker, stream = cell - return (path, marker, stream) != _STREAMED_MESSAGES_OVERFLOW +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: @@ -695,8 +694,7 @@ def _burst_bodies(model: str, round_name: str) -> tuple[_Call, ...]: prompt: Final = f"{marker} {round_name} {path} stream={stream} {tag}" return path, build(model, prompt, stream), marker == _OVERFLOW_MARKER, tag - cells: Final = product(_BUILDERS, (_HAPPY_MARKER, _OVERFLOW_MARKER), (False, True)) - return tuple(call(cell) for cell in cells if _logs_a_spend_row(cell)) + return tuple(call(cell) for cell in product(_BUILDERS, (_HAPPY_MARKER, _OVERFLOW_MARKER), (False, True))) def _tagged_rows(tag: str) -> list[dict[str, JsonValue]]: @@ -732,7 +730,8 @@ def _assert_served(outcomes: tuple[_Outcome, ...]) -> None: continue assert _GENERIC in text, (path, status, text) assert status in (200, 400), (path, status, text) - assert _single_tagged_status(tag) == "failure", (path, tag) + 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: