From 1d8a642e0683e13be122c532436e0919d6c540f5 Mon Sep 17 00:00:00 2001 From: heathriel Date: Wed, 22 Jul 2026 08:41:01 -0700 Subject: [PATCH 01/19] fix(fireworks_ai): support router slugs via routers/ prefix Bare fireworks_ai/ only resolved to accounts/fireworks/models/, so Fireworks routers (served at accounts/fireworks/routers/, e.g. glm-latest and firerouter) could not be reached without passing the full resource id. Add a shared resolve_fireworks_resource_name helper that maps an explicit routers/ or models/ segment to the right resource path, keeps the existing -fast router heuristic, and defaults bare slugs to models/ for backward compatibility. Wire it into both the chat and text-completion transforms, which had drifted (completion lacked router handling entirely) --- .../llms/fireworks_ai/chat/transformation.py | 18 ++++---- litellm/llms/fireworks_ai/common_utils.py | 11 +++++ .../fireworks_ai/completion/transformation.py | 7 +-- .../test_fireworks_ai_chat_transformation.py | 43 ++++++++++++++++++ ..._fireworks_ai_completion_transformation.py | 34 ++++++++++++++ .../test_fireworks_ai_common_utils.py | 45 +++++++++++++++++++ type-discipline-budget.json | 2 +- 7 files changed, 146 insertions(+), 14 deletions(-) create mode 100644 tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py create mode 100644 tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index a796aa47b70..26f0caefacd 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -39,7 +39,11 @@ from ...openai.chat.gpt_transformation import ( OpenAIChatCompletionStreamingHandler, OpenAIGPTConfig, ) -from ..common_utils import FireworksAIException, FireworksAIMixin +from ..common_utils import ( + FireworksAIException, + FireworksAIMixin, + resolve_fireworks_resource_name, +) def _extract_fireworks_hidden_params(payload: dict) -> dict: @@ -459,12 +463,10 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): litellm_params: dict, headers: dict, ) -> dict: - if not model.startswith("accounts/") and "#" not in model: - if model.endswith("-fast"): - model = f"accounts/fireworks/routers/{model}" - else: - model = f"accounts/fireworks/models/{model}" - messages = self._transform_messages_helper(messages=messages, model=model, litellm_params=litellm_params) + resolved_model: Final = resolve_fireworks_resource_name(model) + messages = self._transform_messages_helper( + messages=messages, model=resolved_model, litellm_params=litellm_params + ) if "tools" in optional_params and optional_params["tools"] is not None: tools: Final = self._transform_tools(tools=optional_params["tools"]) optional_params["tools"] = tools @@ -478,7 +480,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): "include_usage": True, } return super().transform_request( - model=model, + model=resolved_model, messages=messages, optional_params=optional_params, litellm_params=litellm_params, diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index 143dd151027..e07e7a26f9e 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -29,6 +29,17 @@ def get_fireworks_session_id(litellm_params: dict) -> str | None: return None +def resolve_fireworks_resource_name(model: str) -> str: + stripped: Final = model.removeprefix("fireworks_ai/") + if stripped.startswith("accounts/") or "#" in stripped: + return stripped + if stripped.startswith(("routers/", "models/")): + return f"accounts/fireworks/{stripped}" + if stripped.endswith("-fast"): + return f"accounts/fireworks/routers/{stripped}" + return f"accounts/fireworks/models/{stripped}" + + class FireworksAIMixin: """ Common Base Config functions across Fireworks AI Endpoints diff --git a/litellm/llms/fireworks_ai/completion/transformation.py b/litellm/llms/fireworks_ai/completion/transformation.py index c141e097d3a..c460510f39c 100644 --- a/litellm/llms/fireworks_ai/completion/transformation.py +++ b/litellm/llms/fireworks_ai/completion/transformation.py @@ -4,7 +4,7 @@ from litellm.types.llms.openai import AllMessageValues, OpenAITextCompletionUser from ...base_llm.completion.transformation import BaseTextCompletionConfig from ...openai.completion.utils import _transform_prompt -from ..common_utils import FireworksAIMixin +from ..common_utils import FireworksAIMixin, resolve_fireworks_resource_name class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig): @@ -50,11 +50,8 @@ class FireworksAITextCompletionConfig(FireworksAIMixin, BaseTextCompletionConfig ) -> dict: prompt: Final = _transform_prompt(messages=messages) - if not model.startswith("accounts/") and "#" not in model: - model = f"accounts/fireworks/models/{model}" - data: Final = { - "model": model, + "model": resolve_fireworks_resource_name(model), "prompt": prompt, **optional_params, } diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 94945ed4bfb..87908ef60c3 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -1282,3 +1282,46 @@ def test_streaming_surfaces_fireworks_response_fields(): assert surfaced["fireworks_raw_outputs"] == [raw_output] assert surfaced["fireworks_perf_metrics"] == {"prompt-tokens": 5} assert surfaced["fireworks_prompt_token_ids"] == [1, 2, 3] + + +def test_transform_request_routes_router_slug(): + config = FireworksAIConfig() + + data = config.transform_request( + model="routers/glm-latest", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert data["model"] == "accounts/fireworks/routers/glm-latest" + + +def test_transform_request_bare_slug_stays_model(): + config = FireworksAIConfig() + + data = config.transform_request( + model="glm-4p6", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert data["model"] == "accounts/fireworks/models/glm-4p6" + + +def test_transform_request_direct_route_passthrough(): + config = FireworksAIConfig() + model = "accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c" + + data = config.transform_request( + model=model, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert data["model"] == model diff --git a/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py b/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py new file mode 100644 index 00000000000..996f1fd975b --- /dev/null +++ b/tests/test_litellm/llms/fireworks_ai/completion/test_fireworks_ai_completion_transformation.py @@ -0,0 +1,34 @@ +import os +import sys + +sys.path.insert(0, os.path.abspath("../../../../..")) + +from litellm.llms.fireworks_ai.completion.transformation import ( + FireworksAITextCompletionConfig, +) + + +def test_transform_text_completion_request_routes_router_slug(): + config = FireworksAITextCompletionConfig() + + data = config.transform_text_completion_request( + model="routers/glm-latest", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + headers={}, + ) + + assert data["model"] == "accounts/fireworks/routers/glm-latest" + + +def test_transform_text_completion_request_bare_slug_stays_model(): + config = FireworksAITextCompletionConfig() + + data = config.transform_text_completion_request( + model="glm-4p6", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + headers={}, + ) + + assert data["model"] == "accounts/fireworks/models/glm-4p6" diff --git a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py new file mode 100644 index 00000000000..4af395baf41 --- /dev/null +++ b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_common_utils.py @@ -0,0 +1,45 @@ +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_name + + +@pytest.mark.parametrize( + "model, expected", + [ + ("routers/glm-latest", "accounts/fireworks/routers/glm-latest"), + ("routers/firerouter", "accounts/fireworks/routers/firerouter"), + ("fireworks_ai/routers/glm-latest", "accounts/fireworks/routers/glm-latest"), + ("models/glm-4p6", "accounts/fireworks/models/glm-4p6"), + ("fireworks_ai/models/glm-4p6", "accounts/fireworks/models/glm-4p6"), + ("glm-4p6", "accounts/fireworks/models/glm-4p6"), + ("fireworks_ai/glm-4p6", "accounts/fireworks/models/glm-4p6"), + ("kimi-k2p6-fast", "accounts/fireworks/routers/kimi-k2p6-fast"), + ( + "accounts/fireworks/routers/glm-latest", + "accounts/fireworks/routers/glm-latest", + ), + ( + "accounts/fireworks/models/glm-4p6", + "accounts/fireworks/models/glm-4p6", + ), + ( + "fireworks_ai/accounts/fireworks/routers/glm-latest", + "accounts/fireworks/routers/glm-latest", + ), + ( + "accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c", + "accounts/fireworks/models/qwen2p5-coder-7b#accounts/gitlab/deployments/2fb7764c", + ), + ( + "glm-4p6#accounts/gitlab/deployments/2fb7764c", + "glm-4p6#accounts/gitlab/deployments/2fb7764c", + ), + ], +) +def test_resolve_fireworks_resource_name(model, expected): + assert resolve_fireworks_resource_name(model) == expected diff --git a/type-discipline-budget.json b/type-discipline-budget.json index ab8198304bb..d9038e20df9 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -30,6 +30,6 @@ "limit": 16783 }, "LIT011": { - "limit": 5602 + "limit": 5599 } } From c99a1ab0d7978a85724fb81c94ab66e704ded309 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Tue, 11 Aug 2026 22:10:17 -0400 Subject: [PATCH 02/19] fix(bedrock): resolve the managed-batch output bucket on the model-routed and cost-poller paths get_configured_s3_bucket_name accepts the output bucket only from the immutable _litellm_internal_model_credentials snapshot or AWS_S3_BUCKET_NAME. That refusal to read litellm_params is deliberate: the bucket is what validate_managed_cloud_file_id checks a file id against, so trusting a request-supplied value would let a caller redirect reads to a bucket of their choosing Two live entry points reach the Bedrock file-content transformation without ever building that snapshot. The managed-files pre-call hook sets data["model"] for any id carrying llm_output_file_id, which is every batch output, so get_file_content always takes the model-routed branch; that branch called llm_router.afile_content directly, and managed_files_obj.afile_content, the only caller that built the snapshot, is therefore unreachable for batch output. CheckBatchCost spread the deployment credentials as plain kwargs, and get_litellm_params does not carry s3_bucket_name across (gcs_bucket_name is listed for exactly this reason, its S3 counterpart is not), so the poller lost the bucket the same way The result was that every completed Bedrock managed batch failed files.content with "S3 bucket_name is required" and never had its cost tracked, leaving the row to be re-polled every cycle. Both paths now resolve the deployment credentials and pass the same MappingProxyType snapshot the managed-files hook already builds --- .../proxy/common_utils/check_batch_cost.py | 2 + .../openai_files_endpoints/files_endpoints.py | 8 ++ .../proxy_unit_tests/test_check_batch_cost.py | 96 +++++++++++++++++++ .../test_files_endpoint.py | 90 +++++++++++++++++ 4 files changed, 196 insertions(+) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 6fe37f0aacb..cfe60a79eed 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -3,6 +3,7 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if t """ from datetime import datetime, timedelta, timezone +from types import MappingProxyType from typing import TYPE_CHECKING, Any, Dict, Final, List, Optional, Tuple from litellm._logging import verbose_proxy_logger @@ -537,6 +538,7 @@ class CheckBatchCost: credentials = self.llm_router.get_deployment_credentials_with_provider(model_id) or {} _file_content = await afile_content( file_id=raw_output_file_id, + _litellm_internal_model_credentials=MappingProxyType(dict(credentials)), **credentials, ) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index d2432ea3729..1cbed2a68f9 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -7,6 +7,7 @@ import asyncio import traceback +from types import MappingProxyType from typing import Any, BinaryIO, Final, cast, get_args import httpx @@ -706,11 +707,18 @@ async def get_file_content( model: Final = cast(str | None, data.get("model")) if model: + deployment_credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model) + trusted_model_credentials: Final = ( + {"_litellm_internal_model_credentials": MappingProxyType(dict(deployment_credentials))} + if deployment_credentials is not None + else {} + ) response = await llm_router.afile_content( **{ "model": model, "file_id": file_id, **data, + **trusted_model_credentials, } ) diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index fa274324fd6..20390159665 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -312,6 +312,102 @@ class TestCheckBatchCost: ), "update() must NOT include batch_processed when column is absent" assert update_data["status"] == "complete" + @pytest.mark.asyncio + async def test_output_fetch_passes_deployment_credentials_as_trusted_snapshot( + self, check_batch_cost_instance, mock_prisma_client, mock_llm_router + ): + """Bedrock resolves the output bucket ONLY from the immutable snapshot kwarg. + + Spreading the credentials as plain kwargs is not enough: get_litellm_params drops + s3_bucket_name, so without _litellm_internal_model_credentials the cost poller + cannot read the output file and every completed Bedrock batch stays unbilled. + """ + from types import MappingProxyType + from unittest.mock import patch + + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + + mock_job = MagicMock() + mock_job.id = "job-bedrock-1" + mock_job.unified_object_id = "dW5pZmllZF9iYXRjaF9pZA==" + mock_job.created_by = "user-1" + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job]) + + mock_response = MagicMock() + mock_response.status = "completed" + mock_response.output_file_id = "file-output-123" + mock_response.model_dump_json.return_value = '{"id":"batch-1","status":"completed"}' + + mock_llm_router.aretrieve_batch = AsyncMock(return_value=mock_response) + mock_llm_router.get_deployment_credentials_with_provider = MagicMock( + return_value={ + "custom_llm_provider": "bedrock", + "s3_bucket_name": "configured-batch-bucket", + "aws_region_name": "us-east-1", + } + ) + + mock_deployment = MagicMock() + mock_deployment.litellm_params.custom_llm_provider = "bedrock" + mock_deployment.litellm_params.model = "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0" + mock_deployment.model_info.model_dump.return_value = {} + mock_llm_router.get_deployment = MagicMock(return_value=mock_deployment) + + mock_file_content = MagicMock() + mock_file_content.content = b'{"recordId":"req-1"}' + + decoded_id = "llm_model_id,model-123;llm_batch_id,batch-456;" + + with ( + patch( + "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id", + side_effect=[decoded_id, None], + ), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id", + return_value="model-123", + ), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id", + return_value="batch-456", + ), + patch( + "litellm.files.main.afile_content", + new_callable=AsyncMock, + return_value=mock_file_content, + ) as mock_afile_content, + patch( + "litellm.batches.batch_utils._get_file_content_as_dictionary", + return_value=[{"recordId": "req-1"}], + ), + patch( + "litellm.batches.batch_utils.calculate_batch_cost_and_usage", + new_callable=AsyncMock, + return_value=(0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["claude-haiku-4-5"]), + ), + patch( + "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", + return_value=("anthropic.claude-haiku-4-5-20251001-v1:0", "bedrock", None, None), + ), + patch("litellm.litellm_core_utils.litellm_logging.Logging") as mock_logging_cls, + ): + mock_logging_obj = MagicMock() + mock_logging_obj.async_success_handler = AsyncMock() + mock_logging_cls.return_value = mock_logging_obj + + await check_batch_cost_instance.check_batch_cost() + + mock_afile_content.assert_awaited() + passed_kwargs = mock_afile_content.await_args[1] + snapshot = passed_kwargs.get("_litellm_internal_model_credentials") + assert snapshot is not None, "cost poller must pass the trusted credential snapshot" + assert isinstance( + snapshot, MappingProxyType + ), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name" + assert snapshot["s3_bucket_name"] == "configured-batch-bucket" + @pytest.mark.asyncio async def test_primary_path_completion_update_includes_batch_processed( self, check_batch_cost_instance, mock_prisma_client, mock_llm_router diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index f27c8dfd2f4..0c26e5f7695 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -3149,6 +3149,96 @@ def test_require_managed_files_rejects_raw_provider_file_id( mock_call.assert_not_called() +def test_get_file_content_model_routed_attaches_trusted_model_credentials(monkeypatch): + """A managed batch output id routes by model, and that branch must build the snapshot. + + The managed-files pre-call hook sets data["model"] for any id carrying + llm_output_file_id, so batch output retrieval always takes the model-routed branch + and never reaches managed_files_obj.afile_content. Bedrock resolves its output + bucket only from _litellm_internal_model_credentials, so without the snapshot every + Bedrock batch output retrieval fails with "S3 bucket_name is required". + """ + import base64 + from types import MappingProxyType + + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + from litellm.types.utils import SpecialEnums + + router = Router( + model_list=[ + { + "model_name": "anthropic.batch.claude-4.5-haiku", + "litellm_params": { + "model": "bedrock/anthropic.claude-haiku-4-5-20251001-v1:0", + "aws_region_name": "us-east-1", + "s3_bucket_name": "configured-batch-bucket", + }, + "model_info": {"id": "bedrock-batch-deployment-id"}, + } + ] + ) + + from unittest.mock import MagicMock + + managed_file_row = MagicMock() + managed_file_row.created_by = "test-user" + managed_file_row.team_id = None + managed_file_row.storage_backend = None + managed_file_row.storage_url = None + prisma_stub = MagicMock() + prisma_stub.db.litellm_managedfiletable.find_first = AsyncMock(return_value=managed_file_row) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_stub) + setup_proxy_logging_object(monkeypatch, router) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) + + captured_kwargs: dict = {} + + async def _mock_router_afile_content(**kwargs): + captured_kwargs.update(kwargs) + return HttpxBinaryResponseContent( + response=httpx.Response( + status_code=200, + content=b'{"recordId":"req-1"}', + headers={"content-type": "application/octet-stream"}, + ) + ) + + monkeypatch.setattr(router, "afile_content", _mock_router_afile_content) + + unified_id = SpecialEnums.LITELLM_MANAGED_FILE_COMPLETE_STR.value.format( + "application/jsonl", + "unified-output-id", + "anthropic.batch.claude-4.5-haiku", + "llm_output_file_id,s3://configured-batch-bucket/out/batch.jsonl", + "bedrock-batch-deployment-id", + ) + encoded_id = base64.urlsafe_b64encode(unified_id.encode()).decode().rstrip("=") + + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + api_key="test-key", + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="test-user", + ) + try: + response = client.get( + f"/v1/files/{encoded_id}/content", + headers={"Authorization": "Bearer test-key", "custom-llm-provider": "bedrock"}, + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + assert response.status_code == 200, response.text + snapshot = captured_kwargs.get("_litellm_internal_model_credentials") + assert snapshot is not None, "model-routed branch must attach the trusted credential snapshot" + assert isinstance( + snapshot, MappingProxyType + ), "snapshot must be a MappingProxyType; a plain dict is rejected by get_configured_s3_bucket_name" + assert snapshot["s3_bucket_name"] == "configured-batch-bucket" + + def _unified_managed_file_id() -> str: import base64 From 460f0d29a95c25091fd375cd8f0f76297525ab78 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Wed, 12 Aug 2026 03:44:12 -0400 Subject: [PATCH 03/19] test(files): capture routed retrieval calls immutably The mock merged every call into one shared dict, so a second routed retrieval would overwrite the first and the assertions would still pass. Keep one frozen snapshot per call and assert exactly one call, which also makes an unintended second retrieval a failure rather than something the merge hides --- .../proxy/openai_files_endpoint/test_files_endpoint.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index 0c26e5f7695..e363a266688 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -3194,10 +3194,12 @@ def test_get_file_content_model_routed_attaches_trusted_model_credentials(monkey monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None) monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", router) - captured_kwargs: dict = {} + # One frozen snapshot per call rather than one dict merged across calls, so a second + # invocation is visible instead of silently overwriting the first. + calls: list[MappingProxyType] = [] async def _mock_router_afile_content(**kwargs): - captured_kwargs.update(kwargs) + calls.append(MappingProxyType(dict(kwargs))) return HttpxBinaryResponseContent( response=httpx.Response( status_code=200, @@ -3231,7 +3233,8 @@ def test_get_file_content_model_routed_attaches_trusted_model_credentials(monkey app.dependency_overrides.pop(ps.user_api_key_auth, None) assert response.status_code == 200, response.text - snapshot = captured_kwargs.get("_litellm_internal_model_credentials") + assert len(calls) == 1, f"expected exactly one routed retrieval, got {len(calls)}" + snapshot = calls[0].get("_litellm_internal_model_credentials") assert snapshot is not None, "model-routed branch must attach the trusted credential snapshot" assert isinstance( snapshot, MappingProxyType From 60fe4e464cc56847735d9c3d3889717f51bee371 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Fri, 14 Aug 2026 00:58:31 -0400 Subject: [PATCH 04/19] fix(bedrock): resolve the managed-batch output bucket on the inline accounting path too A third path reads a completed batch's output file, and it could not resolve the bucket either. When cost is accounted from the retrieve itself rather than from the poller, the batch success handler calls _handle_completed_batch, which fetches the output file through _extract_file_access_credentials. That helper forwarded a whitelist covering Azure and Vertex, gcs_bucket_name included, but nothing for Bedrock, and retrieve_batch built its litellm_params through get_litellm_params, whose fixed signature drops the trusted credential snapshot. So the snapshot never reached the file read and it failed with "S3 bucket_name is required" for a bucket the deployment had configured, leaving the batch's cost unrecorded. Adding s3_bucket_name to that whitelist would not have worked. The Bedrock file config deliberately resolves the bucket only from the immutable server-side snapshot or the environment, never from a request param, because the bucket is what managed file ids are validated against. The snapshot is therefore what has to flow, exactly as it already does for the model-routed and cost-poller paths. retrieve_batch now re-adds the snapshot after get_litellm_params, the same way the file operations already do, the whitelist forwards it, and the proxy attaches it for router-routed managed batches from the deployment behind the unified id. Verified against a live proxy reading a real completed Bedrock batch: the cost row appears within seconds of the retrieve carrying the batch's real spend and usage, where before the read raised and no row was written. Resolving those credentials is best effort. A batch whose deployment no longer resolves, which happens when a model group is removed while batches are in flight, still serves its status instead of failing the request on the lookup. This matters for the OSS and polling-disabled configurations, where the retrieve path is the only thing that accounts for a batch at all. --- litellm/batches/batch_utils.py | 1 + litellm/batches/main.py | 2 + litellm/proxy/batches_endpoints/endpoints.py | 8 +++ .../openai_files_endpoints/common_utils.py | 25 +++++++ .../test_litellm/batches/test_batch_utils.py | 14 ++++ tests/test_litellm/batches/test_main.py | 31 +++++++++ .../proxy/batches_endpoints/test_endpoints.py | 6 +- .../test_files_common_utils.py | 68 +++++++++++++++++++ 8 files changed, 154 insertions(+), 1 deletion(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index e73b887ae0a..f7aa6c50de8 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -309,6 +309,7 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: "bucket_name", "timeout", "max_retries", + "_litellm_internal_model_credentials", ] for key in credential_keys: if key in litellm_params: diff --git a/litellm/batches/main.py b/litellm/batches/main.py index ce52c12818e..bb04d495555 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -22,6 +22,7 @@ from openai.types.batch import BatchRequestCounts import litellm from litellm._logging import verbose_logger +from litellm.files.main import _add_trusted_model_credentials_to_litellm_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler from litellm.llms.azure.batches.handler import AzureBatchesAPI @@ -527,6 +528,7 @@ def retrieve_batch( custom_llm_provider=custom_llm_provider, **kwargs, ) + _add_trusted_model_credentials_to_litellm_params(litellm_params, kwargs) if litellm_logging_obj is not None: litellm_logging_obj.update_from_kwargs( kwargs=kwargs, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index e442cefa360..1301b9327ec 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -25,6 +25,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, apply_team_provider_credentials, decode_model_from_file_id, + add_internal_model_credentials_for_batch, encode_batch_response_ids, encode_file_id_with_model, ensure_batch_response_managed_file_ids, @@ -537,6 +538,13 @@ async def retrieve_batch( detail={"error": "LLM Router not initialized. Ensure models added to proxy."}, ) + if unified_batch_id: + add_internal_model_credentials_for_batch( + data=data, + llm_router=llm_router, + model_id=get_model_id_from_unified_batch_id(unified_batch_id), + ) + response = await llm_router.aretrieve_batch(**data) response._hidden_params["unified_batch_id"] = unified_batch_id if unified_batch_id: diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 56e986c89cf..32676bd1d9f 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -465,6 +465,31 @@ def apply_team_provider_credentials( prepare_data_with_credentials(data=data, credentials=credentials) +def add_internal_model_credentials_for_batch( + data: dict, + llm_router: "Router", + model_id: str | None, +) -> None: + """ + Attach the deployment's immutable server-side credential snapshot to a router-routed + batch call (in-place). + + Cost accounting for a completed batch reads the batch's output file, and the Bedrock + file config resolves its bucket only from this snapshot, never from a request param, + because the bucket is what managed file ids are validated against. Without it that + read fails and the batch's cost is never recorded. + """ + if model_id is None: + return + try: + credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id) + except Exception: # noqa: BLE001 # the snapshot only enables cost accounting; a batch whose deployment no longer resolves must still be retrievable + return + if credentials is None: + return + data["_litellm_internal_model_credentials"] = MappingProxyType(dict(credentials)) + + def prepare_data_with_credentials( data: dict, credentials: dict, diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index 523b512e4cf..cacae3624f3 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -17,6 +17,7 @@ deterministic stand-ins so the arithmetic under test is the only variable. import json import os import sys +from types import MappingProxyType import httpx import pytest @@ -1229,3 +1230,16 @@ async def test_calculate_batch_cost_and_usage_anthropic_end_to_end(): assert cost == pytest.approx(1000 * 3e-6 / 2 + 8000 * 3e-7 / 2 + 2000 * 3.75e-6 / 2 + 200 * 15e-6 / 2) assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (11000, 200, 11200) assert models == ["claude-sonnet-4-5"] + + +def test_extract_credentials_forwards_the_trusted_model_credential_snapshot(): + """Bedrock resolves a batch's output bucket only from the immutable server-side + snapshot, never from a request param, so cost accounting on the retrieve path cannot + read the output file unless this key is forwarded. Without it the accounting raises + "S3 bucket_name is required" for a bucket the deployment has configured, and the + batch's cost is never recorded.""" + snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket", "aws_region_name": "us-east-1"}) + + credentials = bu._extract_file_access_credentials({"_litellm_internal_model_credentials": snapshot}) + + assert credentials["_litellm_internal_model_credentials"] is snapshot diff --git a/tests/test_litellm/batches/test_main.py b/tests/test_litellm/batches/test_main.py index 1f7a91a5511..17e9ee29d4d 100644 --- a/tests/test_litellm/batches/test_main.py +++ b/tests/test_litellm/batches/test_main.py @@ -28,6 +28,7 @@ import sys from contextlib import ExitStack from dataclasses import dataclass from typing import Any, Dict +from types import MappingProxyType from unittest.mock import MagicMock, patch import pytest @@ -742,3 +743,33 @@ def test_resolve_timeout__httpx_timeout_returns_float_read(): resolved = bm._resolve_timeout(_params(timeout=t), {}, "openai") assert isinstance(resolved, float) assert resolved == 99.0 + + +def test_retrieve__forwards_trusted_model_credentials_into_litellm_params(seams): + """The batch's cost is computed by reading its output file after the retrieve, and + Bedrock resolves that bucket only from this immutable snapshot. get_litellm_params has + a fixed signature that drops it, so without re-adding it here the snapshot never + reaches the logging object and cost accounting fails on a bucket that is configured.""" + snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket"}) + logging_obj = MagicMock() + + bm.retrieve_batch( + batch_id="batch-1", + custom_llm_provider="openai", + litellm_logging_obj=logging_obj, + _litellm_internal_model_credentials=snapshot, + ) + + litellm_params = logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"] + assert litellm_params["_litellm_internal_model_credentials"] is snapshot + + +def test_retrieve__omits_trusted_model_credentials_when_not_supplied(seams): + """A retrieve with no snapshot must not invent an empty one, which would read as a + configured bucket of nothing.""" + logging_obj = MagicMock() + + bm.retrieve_batch(batch_id="batch-1", custom_llm_provider="openai", litellm_logging_obj=logging_obj) + + litellm_params = logging_obj.update_from_kwargs.call_args.kwargs["litellm_params"] + assert "_litellm_internal_model_credentials" not in litellm_params diff --git a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py index a80c19f0708..aa5c63280b8 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_endpoints.py @@ -1138,7 +1138,11 @@ async def test_retrieve__unified_batch_id_routes_to_router(retrieve_harness): # DISPATCH - router fired, direct litellm did not. assert retrieve_harness.router_aretrieve.call_count == 1 retrieve_harness.litellm_aretrieve.assert_not_called() - retrieve_harness.creds_resolver.assert_not_called() + + # Credentials are resolved for the deployment behind the unified id so the batch's + # output file can be read for cost accounting. This id resolves to nothing here, and + # the retrieve must still serve the batch rather than fail on the lookup. + retrieve_harness.creds_resolver.assert_called_once_with(model_id="gpt-4o-mini") # router receives the (still-encoded) batch id verbatim - this layer does # not decode it for the unified path. diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index 4a021627c3e..a39f0c5f010 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -95,3 +95,71 @@ def test_apply_unified_file_ids_swaps_all_three_ids(): "unified-out", "unified-err", ) + + +# =========================================================================== # +# add_internal_model_credentials_for_batch - the snapshot that lets a completed +# batch's output file be read, and therefore its cost be recorded +# =========================================================================== # + + +def test_add_internal_model_credentials_attaches_an_immutable_snapshot(): + """Cost accounting for a completed batch reads its output file, and Bedrock resolves + that bucket only from this snapshot. It must be immutable so nothing downstream can + redirect the bucket that managed file ids are validated against.""" + from litellm.proxy.openai_files_endpoints.common_utils import ( + add_internal_model_credentials_for_batch, + ) + + router = MagicMock() + router.get_deployment_credentials_with_provider = MagicMock( + return_value={"s3_bucket_name": "configured-bucket", "aws_region_name": "us-east-1"} + ) + data = {"batch_id": "unified-batch-id"} + + add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id="deployment-1") + + snapshot = data["_litellm_internal_model_credentials"] + assert snapshot["s3_bucket_name"] == "configured-bucket" + assert isinstance(snapshot, MappingProxyType) + with pytest.raises(TypeError): + snapshot["s3_bucket_name"] = "attacker-bucket" + router.get_deployment_credentials_with_provider.assert_called_once_with(model_id="deployment-1") + + +@pytest.mark.parametrize( + "model_id, credentials", + [(None, {"s3_bucket_name": "b"}), ("deployment-1", None)], + ids=["no-model-id", "deployment-has-no-credentials"], +) +def test_add_internal_model_credentials_is_a_noop_without_a_resolvable_deployment(model_id, credentials): + """An unroutable batch must be left alone rather than given an empty snapshot, which + would look like a configured bucket of nothing.""" + from litellm.proxy.openai_files_endpoints.common_utils import ( + add_internal_model_credentials_for_batch, + ) + + router = MagicMock() + router.get_deployment_credentials_with_provider = MagicMock(return_value=credentials) + data = {"batch_id": "unified-batch-id"} + + add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id=model_id) + + assert "_litellm_internal_model_credentials" not in data + + +def test_add_internal_model_credentials_survives_a_failing_deployment_lookup(): + """The snapshot only enables cost accounting, so a batch whose deployment no longer + resolves, which happens when a model group is removed while batches are in flight, + must still be retrievable rather than failing the request on the lookup.""" + from litellm.proxy.openai_files_endpoints.common_utils import ( + add_internal_model_credentials_for_batch, + ) + + router = MagicMock() + router.get_deployment_credentials_with_provider = MagicMock(side_effect=KeyError("deployment-gone")) + data = {"batch_id": "unified-batch-id"} + + add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id="deployment-gone") + + assert data == {"batch_id": "unified-batch-id"} From d7afc1797cf2c0e2c326bf4f2b379991b92a9498 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Fri, 14 Aug 2026 01:32:59 -0400 Subject: [PATCH 05/19] refactor(batches): share the trusted-credentials helper across both call paths The helper that carries the credential snapshot into litellm_params lived private in files/main.py, and the batch retrieve needed it too. It now sits beside get_litellm_params, which is what it augments, so neither caller reaches into the other's private surface. Typed as Mapping/MutableMapping of object rather than Any, which the strict import rules ban. The file-content route builds the snapshot through the same helper as the batch route instead of assembling a conditional mapping inline, which drops two mutable constructions and leaves one way to attach it. Its name loses the batch suffix now that both routes use it. --- litellm/batches/main.py | 4 ++-- litellm/files/main.py | 16 ++++------------ .../litellm_core_utils/get_litellm_params.py | 18 ++++++++++++++++++ litellm/proxy/batches_endpoints/endpoints.py | 4 ++-- .../openai_files_endpoints/common_utils.py | 2 +- .../openai_files_endpoints/files_endpoints.py | 10 ++-------- .../test_files_common_utils.py | 14 +++++++------- 7 files changed, 36 insertions(+), 32 deletions(-) diff --git a/litellm/batches/main.py b/litellm/batches/main.py index bb04d495555..20d38bbb77f 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -22,7 +22,7 @@ from openai.types.batch import BatchRequestCounts import litellm from litellm._logging import verbose_logger -from litellm.files.main import _add_trusted_model_credentials_to_litellm_params +from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.anthropic.batches.handler import AnthropicBatchesHandler from litellm.llms.azure.batches.handler import AzureBatchesAPI @@ -528,7 +528,7 @@ def retrieve_batch( custom_llm_provider=custom_llm_provider, **kwargs, ) - _add_trusted_model_credentials_to_litellm_params(litellm_params, kwargs) + add_trusted_model_credentials_to_litellm_params(litellm_params, kwargs) if litellm_logging_obj is not None: litellm_logging_obj.update_from_kwargs( kwargs=kwargs, diff --git a/litellm/files/main.py b/litellm/files/main.py index 34421d13761..9a64c78552b 100644 --- a/litellm/files/main.py +++ b/litellm/files/main.py @@ -11,7 +11,6 @@ import time import uuid as uuid_module from collections.abc import Coroutine from functools import partial -from types import MappingProxyType from typing import Any, Final, Literal, cast import httpx @@ -34,6 +33,7 @@ import litellm from litellm import get_secret_str from litellm.files.streaming import FileContentStreamingResponse from litellm.files.types import FileContentProvider, FileContentStreamingResult +from litellm.litellm_core_utils.get_litellm_params import add_trusted_model_credentials_to_litellm_params from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.azure.common_utils import get_azure_credentials @@ -85,14 +85,6 @@ bedrock_files_instance: Final = BedrockFilesHandler() ################################################# -def _add_trusted_model_credentials_to_litellm_params( - litellm_params_dict: dict[str, Any], kwargs: dict[str, Any] -) -> None: - trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials") - if isinstance(trusted_model_credentials, type(MappingProxyType({}))): - litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials - - @client async def acreate_file( file: FileTypes, @@ -372,7 +364,7 @@ def file_retrieve( ) if provider_config is not None: litellm_params_dict: Final = get_litellm_params(**kwargs) - _add_trusted_model_credentials_to_litellm_params( + add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, ) @@ -494,7 +486,7 @@ def file_delete( pass optional_params: Final = GenericLiteLLMParams(**kwargs) litellm_params_dict: Final = get_litellm_params(**kwargs) - _add_trusted_model_credentials_to_litellm_params( + add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, ) @@ -834,7 +826,7 @@ def file_content( try: optional_params: Final = GenericLiteLLMParams(**kwargs) litellm_params_dict: Final = get_litellm_params(**kwargs) - _add_trusted_model_credentials_to_litellm_params( + add_trusted_model_credentials_to_litellm_params( litellm_params_dict=litellm_params_dict, kwargs=kwargs, ) diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index f251ab4d74a..3eb8c163d5c 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -1,3 +1,5 @@ +from collections.abc import Mapping, MutableMapping +from types import MappingProxyType from typing import Final from litellm.llms.openai.data_residency import infer_openai_data_residency @@ -184,3 +186,19 @@ def get_litellm_params( litellm_params[key] = kwargs[key] return litellm_params + + +def add_trusted_model_credentials_to_litellm_params( + litellm_params_dict: MutableMapping[str, object], kwargs: Mapping[str, object] +) -> None: + """ + Carry the immutable server-side credential snapshot into litellm_params. + + get_litellm_params has a fixed signature, so callers that need the snapshot to + survive into the logging object and the downstream file read have to re-add it. Only + a MappingProxyType is accepted, since providers resolve trusted configuration such + as a Bedrock file bucket from it and must not read a request-supplied mapping. + """ + trusted_model_credentials: Final = kwargs.get("_litellm_internal_model_credentials") + if isinstance(trusted_model_credentials, MappingProxyType): + litellm_params_dict["_litellm_internal_model_credentials"] = trusted_model_credentials diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 1301b9327ec..9a6bb054d1f 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -25,7 +25,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, apply_team_provider_credentials, decode_model_from_file_id, - add_internal_model_credentials_for_batch, + add_internal_model_credentials, encode_batch_response_ids, encode_file_id_with_model, ensure_batch_response_managed_file_ids, @@ -539,7 +539,7 @@ async def retrieve_batch( ) if unified_batch_id: - add_internal_model_credentials_for_batch( + add_internal_model_credentials( data=data, llm_router=llm_router, model_id=get_model_id_from_unified_batch_id(unified_batch_id), diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 32676bd1d9f..f2e6fb633e1 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -465,7 +465,7 @@ def apply_team_provider_credentials( prepare_data_with_credentials(data=data, credentials=credentials) -def add_internal_model_credentials_for_batch( +def add_internal_model_credentials( data: dict, llm_router: "Router", model_id: str | None, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 1cbed2a68f9..361b5b920e2 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -7,7 +7,6 @@ import asyncio import traceback -from types import MappingProxyType from typing import Any, BinaryIO, Final, cast, get_args import httpx @@ -44,6 +43,7 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + add_internal_model_credentials, apply_team_provider_credentials, encode_file_id_with_model, extract_file_creation_params, @@ -707,18 +707,12 @@ async def get_file_content( model: Final = cast(str | None, data.get("model")) if model: - deployment_credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model) - trusted_model_credentials: Final = ( - {"_litellm_internal_model_credentials": MappingProxyType(dict(deployment_credentials))} - if deployment_credentials is not None - else {} - ) + add_internal_model_credentials(data=data, llm_router=llm_router, model_id=model) response = await llm_router.afile_content( **{ "model": model, "file_id": file_id, **data, - **trusted_model_credentials, } ) diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py index a39f0c5f010..ad7f5e4725a 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_common_utils.py @@ -98,7 +98,7 @@ def test_apply_unified_file_ids_swaps_all_three_ids(): # =========================================================================== # -# add_internal_model_credentials_for_batch - the snapshot that lets a completed +# add_internal_model_credentials - the snapshot that lets a completed # batch's output file be read, and therefore its cost be recorded # =========================================================================== # @@ -108,7 +108,7 @@ def test_add_internal_model_credentials_attaches_an_immutable_snapshot(): that bucket only from this snapshot. It must be immutable so nothing downstream can redirect the bucket that managed file ids are validated against.""" from litellm.proxy.openai_files_endpoints.common_utils import ( - add_internal_model_credentials_for_batch, + add_internal_model_credentials, ) router = MagicMock() @@ -117,7 +117,7 @@ def test_add_internal_model_credentials_attaches_an_immutable_snapshot(): ) data = {"batch_id": "unified-batch-id"} - add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id="deployment-1") + add_internal_model_credentials(data=data, llm_router=router, model_id="deployment-1") snapshot = data["_litellm_internal_model_credentials"] assert snapshot["s3_bucket_name"] == "configured-bucket" @@ -136,14 +136,14 @@ def test_add_internal_model_credentials_is_a_noop_without_a_resolvable_deploymen """An unroutable batch must be left alone rather than given an empty snapshot, which would look like a configured bucket of nothing.""" from litellm.proxy.openai_files_endpoints.common_utils import ( - add_internal_model_credentials_for_batch, + add_internal_model_credentials, ) router = MagicMock() router.get_deployment_credentials_with_provider = MagicMock(return_value=credentials) data = {"batch_id": "unified-batch-id"} - add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id=model_id) + add_internal_model_credentials(data=data, llm_router=router, model_id=model_id) assert "_litellm_internal_model_credentials" not in data @@ -153,13 +153,13 @@ def test_add_internal_model_credentials_survives_a_failing_deployment_lookup(): resolves, which happens when a model group is removed while batches are in flight, must still be retrievable rather than failing the request on the lookup.""" from litellm.proxy.openai_files_endpoints.common_utils import ( - add_internal_model_credentials_for_batch, + add_internal_model_credentials, ) router = MagicMock() router.get_deployment_credentials_with_provider = MagicMock(side_effect=KeyError("deployment-gone")) data = {"batch_id": "unified-batch-id"} - add_internal_model_credentials_for_batch(data=data, llm_router=router, model_id="deployment-gone") + add_internal_model_credentials(data=data, llm_router=router, model_id="deployment-gone") assert data == {"batch_id": "unified-batch-id"} From eafddaaa12a75ad70d9a1dc4ec90f3f611e8a130 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 17:22:32 -0700 Subject: [PATCH 06/19] feat(scripts): queue heavy gates behind a machine-wide slot lock --- Makefile | 21 +- scripts/gate_slot_lock.py | 173 ++++++++++++ scripts/pre_commit_lint.sh | 10 + scripts/ruff_strict_gate.py | 5 +- scripts/type_check_gate.py | 23 +- scripts/type_discipline_gate.py | 5 +- tests/test_litellm/test_gate_slot_lock.py | 311 +++++++++++++++++++++ tests/test_litellm/test_pre_commit_lint.py | 41 +++ 8 files changed, 573 insertions(+), 16 deletions(-) create mode 100644 scripts/gate_slot_lock.py create mode 100644 tests/test_litellm/test_gate_slot_lock.py diff --git a/Makefile b/Makefile index 94d8c875af5..7fe5d1f8045 100644 --- a/Makefile +++ b/Makefile @@ -8,8 +8,8 @@ lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \ lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \ install-dev install-proxy-dev install-test-deps install-hooks \ - install-helm-unittest check-circular-imports check-import-safety check pre-commit \ - lint-install lint-fetch-base bootstrap + install-helm-unittest check-circular-imports check-import-safety check check-inner pre-commit \ + lint-install lint-fetch-base bootstrap bootstrap-inner # Default target help: @@ -52,10 +52,17 @@ help: @echo " make test-proxy-unit-b - Run proxy_unit_tests (p-z, ~28 files)" @echo " make test-integration - Run integration tests" @echo " make test-unit-helm - Run helm unit tests" + @echo "" + @echo "Heavy targets (check, bootstrap, lint) queue for LITELLM_GATE_SLOTS machine-wide" + @echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine." UV := uv UV_RUN := $(UV) run --no-sync +# Machine-wide slot queue for the heavy targets below; python3 + stdlib only, so +# it runs before any venv exists. See scripts/gate_slot_lock.py. +GATE_SLOT_LOCK := python3 scripts/gate_slot_lock.py + LINT_DEP_INSTALL ?= install-dev LINT_E2E_DEP_INSTALL ?= lint-install LINT_DEP_BASE ?= lint-fetch-base @@ -74,6 +81,9 @@ install-dev: $(UV) sync --inexact --frozen bootstrap: + @$(GATE_SLOT_LOCK) $(MAKE) bootstrap-inner + +bootstrap-inner: $(UV) sync --inexact --frozen --extra proxy --group proxy-dev --group e2e-dev $(UV_RUN) python scripts/prisma_generate_if_needed.py cd ui/litellm-dashboard && ../../scripts/with_dashboard_node.sh npm install --no-audit --no-fund @@ -230,7 +240,7 @@ check-import-safety: $(LINT_DEP_INSTALL) # base fetch) runs once up front; the checks themselves are independent, so a sub-make # fans them out with -j and the fast ones finish under basedpyright's shadow. lint: lint-install lint-fetch-base - $(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks + $(GATE_SLOT_LOCK) $(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety @@ -244,7 +254,10 @@ lint-dev: lint-format-changed check-circular-imports check-import-safety # test-linting.yml (Python), test-litellm-ui-build.yml's frontend-lint (dashboard), and # check-ui-api-types.yml (API-type drift), skipping any whose files aren't in scope. # Not auto-installed as a git hook so it never slows an unrelated human commit. -check: bootstrap +check: + @$(GATE_SLOT_LOCK) $(MAKE) check-inner + +check-inner: bootstrap ./scripts/pre_commit_lint.sh pre-commit: diff --git a/scripts/gate_slot_lock.py b/scripts/gate_slot_lock.py new file mode 100644 index 00000000000..e7bd945ad65 --- /dev/null +++ b/scripts/gate_slot_lock.py @@ -0,0 +1,173 @@ +#!/usr/bin/env python3 +"""Machine-wide slot lock for this repo's heavy entrypoints. + +`make check`, `make bootstrap`, `make lint`, and the standalone budget gates +(scripts/ruff_strict_gate.py, scripts/type_discipline_gate.py, +scripts/type_check_gate.py) each hold one of N machine-wide slots while they +run, so however many sessions and worktrees share one machine, at most N of +them execute a basedpyright/pytest/prettier storm at a time instead of all +thrashing it at once. Slots are fcntl.flock files (macOS ships no flock(1) +binary, hence python3 + stdlib only, runnable before any venv exists) under a +per-user cache directory shared by every worktree and session: +~/.cache/litellm/gate-slots by default, $LITELLM_GATE_SLOT_DIR to override. +A holder's lock dies with its process, so a crash leaves nothing to clean up. + +$LITELLM_GATE_SLOTS sets the slot count (default 2); 0 disables locking. +Waiting is a blocking flock on a turnstile file plus a slow poll of the slots, +so contenders queue roughly first-come-first-served without busy-spinning. +A process that acquired (or deliberately skipped) a slot exports +LITELLM_GATE_SLOT_HELD, and nested acquisitions under that marker are no-ops, +so `make check` invoking the gates internally can never deadlock against +itself. Any filesystem error fails open and the command runs unlocked: the +lock is a courtesy to the machine, never a gate that may break a build (CI +runs one job per machine, so there it only ever takes the instant path). + +CLI: python3 scripts/gate_slot_lock.py [args...] +""" + +from __future__ import annotations + +import contextlib +import fcntl +import os +import subprocess +import sys +import time +from pathlib import Path +from typing import IO, TYPE_CHECKING, Final + +if TYPE_CHECKING: + from collections.abc import Iterator + +HELD_MARKER_ENV: Final = "LITELLM_GATE_SLOT_HELD" +SLOT_COUNT_ENV: Final = "LITELLM_GATE_SLOTS" +SLOT_DIR_ENV: Final = "LITELLM_GATE_SLOT_DIR" +DEFAULT_SLOT_COUNT: Final = 2 +POLL_SECONDS: Final = 2.0 + + +def _slot_dir() -> Path: + override: Final = os.environ.get(SLOT_DIR_ENV) + return Path(override) if override else Path.home() / ".cache" / "litellm" / "gate-slots" + + +def _slot_count() -> int: + raw: Final = os.environ.get(SLOT_COUNT_ENV) + if not raw: + return DEFAULT_SLOT_COUNT + try: + return int(raw) + except ValueError: + print( + f"gate_slot_lock: ignoring non-integer {SLOT_COUNT_ENV}={raw!r}; " + f"using {DEFAULT_SLOT_COUNT} slots", + file=sys.stderr, + ) + return DEFAULT_SLOT_COUNT + + +def _try_slot(directory: Path, index: int) -> IO[bytes] | None: + handle: Final = (directory / f"slot-{index}.lock").open("wb") + try: + fcntl.flock(handle, fcntl.LOCK_EX | fcntl.LOCK_NB) + except BlockingIOError: + handle.close() + return None + except OSError: + handle.close() + raise + return handle + + +def _wait_for_slot(directory: Path, count: int) -> IO[bytes]: + print( + f"gate_slot_lock: all {count} machine-wide slots are busy; queueing " + f"(set {SLOT_COUNT_ENV}=0 to disable)", + file=sys.stderr, + flush=True, + ) + with (directory / "turnstile.lock").open("wb") as turnstile: + fcntl.flock(turnstile, fcntl.LOCK_EX) + while True: + for index in range(count): + held = _try_slot(directory, index) + if held is not None: + return held + time.sleep(POLL_SECONDS) + + +def _locked_handle(count: int) -> IO[bytes]: + directory: Final = _slot_dir() + directory.mkdir(parents=True, exist_ok=True) + for index in range(count): + immediate = _try_slot(directory, index) + if immediate is not None: + return immediate + return _wait_for_slot(directory, count) + + +def acquire_slot() -> IO[bytes] | None: + """Hold a machine-wide slot for the life of the returned handle. + + The caller must keep the handle referenced until the process exits; + dropping it closes the file and releases the slot. Returns None without + locking when this process already runs under a held slot, when locking is + disabled, or when the filesystem refuses to cooperate.""" + if os.environ.get(HELD_MARKER_ENV): + return None + count: Final = _slot_count() + if count <= 0: + os.environ[HELD_MARKER_ENV] = "1" + return None + try: + handle: Final = _locked_handle(count) + except (OSError, RuntimeError) as error: + print(f"gate_slot_lock: locking unavailable ({error}); running unlocked", file=sys.stderr) + os.environ[HELD_MARKER_ENV] = "1" + return None + os.environ[HELD_MARKER_ENV] = "1" + return handle + + +@contextlib.contextmanager +def held_slot() -> Iterator[None]: + """Run the with-block while holding a machine-wide slot (or its no-op forms).""" + prior_marker: Final = os.environ.get(HELD_MARKER_ENV) + handle: Final = acquire_slot() + try: + yield + finally: + if handle is not None: + handle.close() + if not prior_marker: + os.environ.pop(HELD_MARKER_ENV, None) + + +def _wait_ignoring_interrupts(process: subprocess.Popen[bytes]) -> int: + while True: + try: + return process.wait() + except KeyboardInterrupt: + continue + + +def main() -> int: + if len(sys.argv) < 2: + print("usage: gate_slot_lock.py [args...]", file=sys.stderr) + return 2 + try: + held: Final = acquire_slot() + except KeyboardInterrupt: + return 130 + try: + code: Final = _wait_ignoring_interrupts(subprocess.Popen(sys.argv[1:])) + except FileNotFoundError as error: + print(f"gate_slot_lock: {error}", file=sys.stderr) + return 127 + if held is not None: + held.close() + return code if code >= 0 else 128 - code + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index 82498ec10cd..0861172056e 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -24,6 +24,16 @@ set -eu +# Queue for one of the machine-wide heavy-work slots (see scripts/gate_slot_lock.py) +# before anything else, so N parallel `make check` runs across worktrees execute two +# at a time instead of thrashing the machine. The wrapper exports +# LITELLM_GATE_SLOT_HELD, so this re-exec happens exactly once and everything this +# script spawns (make lint, the budget gates) skips its own acquisition. +if [ -z "${LITELLM_GATE_SLOT_HELD:-}" ]; then + script_dir=$(python3 -c 'import os, sys; print(os.path.dirname(os.path.realpath(sys.argv[1])))' "$0") + exec python3 "$script_dir/gate_slot_lock.py" "$0" "$@" +fi + if [ -z "${PRE_COMMIT_LINT_INNER:-}" ]; then log_file=$(git rev-parse --path-format=absolute --git-path pre_commit_lint.log) if : > "$log_file" 2>/dev/null; then diff --git a/scripts/ruff_strict_gate.py b/scripts/ruff_strict_gate.py index 507077ddf25..bf070beeb0f 100644 --- a/scripts/ruff_strict_gate.py +++ b/scripts/ruff_strict_gate.py @@ -215,7 +215,10 @@ def main() -> None: parser.add_argument("--base", default=DEFAULT_BASE) parser.add_argument("--update", action="store_true") args = parser.parse_args() - cmd_update(args.base) if args.update else cmd_check(args.base) + from gate_slot_lock import held_slot + + with held_slot(): + cmd_update(args.base) if args.update else cmd_check(args.base) if __name__ == "__main__": diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index c9f774c6113..763835e6d2e 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -670,16 +670,19 @@ def main() -> None: parser.add_argument("--update", action="store_true") parser.add_argument("--emit-counts-dir", type=Path) args = parser.parse_args() - ensure_typecheck_env() - head = count_basedpyright(run_basedpyright()) - if args.emit_counts_dir is not None: - cmd_emit_counts( - head, args.emit_counts_dir, _run(["git", "rev-parse", "HEAD"]).strip() - ) - elif args.update: - cmd_update(head, args.base) - else: - cmd_check(head, args.base) + from gate_slot_lock import held_slot + + with held_slot(): + ensure_typecheck_env() + head = count_basedpyright(run_basedpyright()) + if args.emit_counts_dir is not None: + cmd_emit_counts( + head, args.emit_counts_dir, _run(["git", "rev-parse", "HEAD"]).strip() + ) + elif args.update: + cmd_update(head, args.base) + else: + cmd_check(head, args.base) if __name__ == "__main__": diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py index f937283d972..5f6474f20bc 100644 --- a/scripts/type_discipline_gate.py +++ b/scripts/type_discipline_gate.py @@ -267,7 +267,10 @@ def main() -> None: parser.add_argument("--base", default=DEFAULT_BASE) parser.add_argument("--update", action="store_true") args = parser.parse_args() - cmd_update(args.base) if args.update else cmd_check(args.base) + from gate_slot_lock import held_slot + + with held_slot(): + cmd_update(args.base) if args.update else cmd_check(args.base) if __name__ == "__main__": diff --git a/tests/test_litellm/test_gate_slot_lock.py b/tests/test_litellm/test_gate_slot_lock.py new file mode 100644 index 00000000000..1cf52ae89f6 --- /dev/null +++ b/tests/test_litellm/test_gate_slot_lock.py @@ -0,0 +1,311 @@ +import fcntl +import importlib.util +import os +import signal +import subprocess +import sys +import time +from collections.abc import Callable, Sequence +from contextlib import suppress +from pathlib import Path + +import pytest + +ROOT = Path(__file__).resolve().parents[2] +HELPER = ROOT / "scripts" / "gate_slot_lock.py" + +_spec = importlib.util.spec_from_file_location("gate_slot_lock", HELPER) +assert _spec is not None and _spec.loader is not None +gate_slot_lock = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(gate_slot_lock) + +START_THEN_WAIT_FOR = ( + "import pathlib, sys, time\n" + "pathlib.Path(sys.argv[1]).touch()\n" + "deadline = time.monotonic() + 20\n" + "while not pathlib.Path(sys.argv[2]).exists():\n" + " if time.monotonic() > deadline:\n" + " sys.exit(3)\n" + " time.sleep(0.05)\n" +) + +TOUCH_TARGET = "import pathlib, sys\npathlib.Path(sys.argv[1]).touch()\n" + +RECORD_INTERVAL = ( + "import sys, time\n" + "with open(sys.argv[1], 'a') as events:\n" + " events.write(f'start {time.monotonic()}\\n')\n" + " events.flush()\n" + " time.sleep(0.6)\n" + " events.write(f'end {time.monotonic()}\\n')\n" + " events.flush()\n" +) + + +def _env(lock_dir: Path, slots: str) -> dict[str, str]: + return { + "PATH": os.environ["PATH"], + "HOME": str(lock_dir.parent), + "LITELLM_GATE_SLOT_DIR": str(lock_dir), + "LITELLM_GATE_SLOTS": slots, + } + + +def _wrapped(payload: Sequence[str]) -> list[str]: + return [sys.executable, str(HELPER), sys.executable, "-c", *payload] + + +def _wait_until(predicate: Callable[[], bool], timeout_seconds: float) -> bool: + deadline = time.monotonic() + timeout_seconds + while time.monotonic() < deadline: + if predicate(): + return True + time.sleep(0.05) + return predicate() + + +def _terminate_group(process: subprocess.Popen[bytes]) -> None: + with suppress(ProcessLookupError, PermissionError): + os.killpg(process.pid, signal.SIGKILL) + + +def _reap(process: subprocess.Popen[bytes]) -> None: + with suppress(subprocess.TimeoutExpired): + process.wait(timeout=10) + if process.poll() is None: + process.kill() + process.wait(timeout=10) + + +def test_six_contenders_never_exceed_two_slots_and_all_complete(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + events_file = tmp_path / "events.log" + env = _env(lock_dir, "2") + procs = [ + subprocess.Popen( + _wrapped([RECORD_INTERVAL, str(events_file)]), + env=env, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + for _ in range(6) + ] + try: + assert [proc.wait(timeout=60) for proc in procs] == [0] * 6 + finally: + for proc in procs: + if proc.poll() is None: + proc.kill() + proc.wait(timeout=10) + events = sorted( + (float(stamp), 1 if kind == "start" else -1) + for kind, stamp in (line.split() for line in events_file.read_text().splitlines()) + ) + assert len(events) == 12 + concurrency_peaks = [] + running = 0 + for _, delta in events: + running += delta + concurrency_peaks.append(running) + assert max(concurrency_peaks) <= 2 + + +def test_two_slots_admit_two_holders_at_once(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + first_started = tmp_path / "first.started" + second_started = tmp_path / "second.started" + env = _env(lock_dir, "2") + first = subprocess.Popen( + _wrapped([START_THEN_WAIT_FOR, str(first_started), str(second_started)]), env=env + ) + second = subprocess.Popen( + _wrapped([START_THEN_WAIT_FOR, str(second_started), str(first_started)]), env=env + ) + assert first.wait(timeout=30) == 0 + assert second.wait(timeout=30) == 0 + + +def test_contender_beyond_capacity_queues_until_the_slot_frees(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + holder_started = tmp_path / "holder.started" + release = tmp_path / "release" + done = tmp_path / "done" + env = _env(lock_dir, "1") + holder = subprocess.Popen( + _wrapped([START_THEN_WAIT_FOR, str(holder_started), str(release)]), env=env + ) + try: + assert _wait_until(holder_started.exists, 10) + contender = subprocess.Popen( + _wrapped([TOUCH_TARGET, str(done)]), + env=env, + stderr=subprocess.PIPE, + ) + try: + time.sleep(1.5) + assert not done.exists() + release.touch() + assert holder.wait(timeout=10) == 0 + assert contender.wait(timeout=30) == 0 + assert done.exists() + assert contender.stderr is not None + assert b"queueing" in contender.stderr.read() + finally: + release.touch() + _reap(contender) + finally: + release.touch() + _reap(holder) + + +def test_nested_wrapping_reenters_instead_of_deadlocking(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + nested = [ + sys.executable, + str(HELPER), + sys.executable, + str(HELPER), + sys.executable, + "-c", + "print('nested ok')", + ] + proc = subprocess.Popen( + nested, + env=_env(lock_dir, "1"), + stdout=subprocess.PIPE, + start_new_session=True, + ) + try: + stdout, _ = proc.communicate(timeout=20) + except subprocess.TimeoutExpired: + _terminate_group(proc) + pytest.fail("nested gate_slot_lock invocations deadlocked") + assert proc.returncode == 0 + assert b"nested ok" in stdout + + +def test_wrapped_command_exit_code_is_propagated(tmp_path: Path) -> None: + proc = subprocess.run( + [sys.executable, str(HELPER), sys.executable, "-c", "raise SystemExit(7)"], + env=_env(tmp_path / "locks", "2"), + ) + assert proc.returncode == 7 + + +def test_missing_command_exits_127_and_no_command_exits_2(tmp_path: Path) -> None: + env = _env(tmp_path / "locks", "2") + missing = subprocess.run( + [sys.executable, str(HELPER), str(tmp_path / "no-such-binary")], + env=env, + capture_output=True, + ) + assert missing.returncode == 127 + bare = subprocess.run([sys.executable, str(HELPER)], env=env, capture_output=True) + assert bare.returncode == 2 + + +def test_wrapped_command_killed_by_signal_maps_to_128_plus_signal(tmp_path: Path) -> None: + proc = subprocess.run( + _wrapped(["import os, signal\nos.kill(os.getpid(), signal.SIGTERM)\n"]), + env=_env(tmp_path / "locks", "2"), + ) + assert proc.returncode == 128 + signal.SIGTERM + + +def test_unusable_lock_dir_fails_open_and_still_runs_the_command(tmp_path: Path) -> None: + blocker = tmp_path / "blocker" + blocker.write_text("") + done = tmp_path / "done" + proc = subprocess.run( + _wrapped([TOUCH_TARGET, str(done)]), + env=_env(blocker / "locks", "2"), + capture_output=True, + ) + assert proc.returncode == 0 + assert done.exists() + assert b"running unlocked" in proc.stderr + + +def test_zero_slots_disables_locking_entirely(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + done = tmp_path / "done" + proc = subprocess.run( + _wrapped([TOUCH_TARGET, str(done)]), + env=_env(lock_dir, "0"), + ) + assert proc.returncode == 0 + assert done.exists() + assert not lock_dir.exists() + + +def test_non_integer_slot_count_warns_and_falls_back_to_default(tmp_path: Path) -> None: + proc = subprocess.run( + [sys.executable, str(HELPER), sys.executable, "-c", "print('ran')"], + env=_env(tmp_path / "locks", "lots"), + capture_output=True, + ) + assert proc.returncode == 0 + assert b"ran" in proc.stdout + assert b"LITELLM_GATE_SLOTS" in proc.stderr + + +def test_killed_holder_releases_its_slot_for_the_next_contender(tmp_path: Path) -> None: + lock_dir = tmp_path / "locks" + holder_started = tmp_path / "holder.started" + never = tmp_path / "never" + env = _env(lock_dir, "1") + holder = subprocess.Popen( + _wrapped([START_THEN_WAIT_FOR, str(holder_started), str(never)]), + env=env, + start_new_session=True, + ) + try: + assert _wait_until(holder_started.exists, 10) + finally: + _terminate_group(holder) + holder.wait(timeout=10) + after = subprocess.run( + [sys.executable, str(HELPER), sys.executable, "-c", "print('freed')"], + env=env, + capture_output=True, + timeout=20, + ) + assert after.returncode == 0 + assert b"freed" in after.stdout + + +def test_acquire_slot_holds_marks_and_releases_in_process( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + lock_dir = tmp_path / "locks" + monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "") + monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir)) + monkeypatch.setenv("LITELLM_GATE_SLOTS", "1") + handle = gate_slot_lock.acquire_slot() + assert handle is not None + assert os.environ["LITELLM_GATE_SLOT_HELD"] == "1" + assert gate_slot_lock.acquire_slot() is None + with (lock_dir / "slot-0.lock").open("wb") as probe: + with pytest.raises(BlockingIOError): + fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB) + handle.close() + fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB) + fcntl.flock(probe, fcntl.LOCK_UN) + + +def test_held_slot_context_manager_releases_on_exit( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + lock_dir = tmp_path / "locks" + monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "") + monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir)) + monkeypatch.setenv("LITELLM_GATE_SLOTS", "1") + with gate_slot_lock.held_slot(): + assert os.environ["LITELLM_GATE_SLOT_HELD"] == "1" + with (lock_dir / "slot-0.lock").open("wb") as probe: + with pytest.raises(BlockingIOError): + fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB) + assert not os.environ.get("LITELLM_GATE_SLOT_HELD") + with (lock_dir / "slot-0.lock").open("wb") as probe: + fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB) + fcntl.flock(probe, fcntl.LOCK_UN) diff --git a/tests/test_litellm/test_pre_commit_lint.py b/tests/test_litellm/test_pre_commit_lint.py index 33baf7474ce..5ea0e79a196 100644 --- a/tests/test_litellm/test_pre_commit_lint.py +++ b/tests/test_litellm/test_pre_commit_lint.py @@ -1,4 +1,5 @@ import os +import shutil import signal import subprocess import time @@ -420,6 +421,46 @@ def test_staged_files_matching_no_check_print_an_explicit_noop_note_and_nonempty assert "skipped: Python lint (make lint) (no litellm/ Python files in scope)" in log +def test_run_queues_through_the_machine_wide_gate_slot_lock(tmp_path: Path) -> None: + repo, bin_dir = _sandbox(tmp_path) + lock_dir = tmp_path / "gate-locks" + proc = _run(repo, bin_dir, {"LITELLM_GATE_SLOT_DIR": str(lock_dir)}) + assert proc.returncode == 0, proc.stdout + proc.stderr + assert (lock_dir / "slot-0.lock").exists() + + +def test_run_under_a_held_slot_skips_reacquiring_the_gate_lock(tmp_path: Path) -> None: + repo, bin_dir = _sandbox(tmp_path) + lock_dir = tmp_path / "gate-locks" + proc = _run( + repo, + bin_dir, + {"LITELLM_GATE_SLOT_DIR": str(lock_dir), "LITELLM_GATE_SLOT_HELD": "1"}, + ) + assert proc.returncode == 0, proc.stdout + proc.stderr + assert not lock_dir.exists() + + +def test_hook_symlink_install_still_resolves_the_slot_lock_helper(tmp_path: Path) -> None: + repo, bin_dir = _sandbox(tmp_path) + scripts_dir = repo / "scripts" + scripts_dir.mkdir() + shutil.copy(SCRIPT, scripts_dir / "pre_commit_lint.sh") + shutil.copy(SCRIPT.parent / "gate_slot_lock.py", scripts_dir / "gate_slot_lock.py") + (repo / ".git" / "hooks" / "pre-commit").symlink_to(Path("../../scripts/pre_commit_lint.sh")) + lock_dir = tmp_path / "gate-locks" + proc = subprocess.run( + ["git", "-c", "user.email=t@t", "-c", "user.name=t", "commit", "-qm", "hooked"], + cwd=repo, + capture_output=True, + text=True, + env=_env(repo, bin_dir, {"LITELLM_GATE_SLOT_DIR": str(lock_dir)}), + timeout=120, + ) + assert proc.returncode == 0, proc.stdout + proc.stderr + assert (lock_dir / "slot-0.lock").exists() + + def test_failing_run_ends_with_a_fail_verdict(tmp_path: Path) -> None: repo, bin_dir = _sandbox(tmp_path) proc = _run(repo, bin_dir, {"STUB_FAIL": "make-lint"}) From 17f5c909f06b16ad554b6aa005458bff2a323b78 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Fri, 14 Aug 2026 21:17:14 -0700 Subject: [PATCH 07/19] fix(make): acquire the gate slot before lint setup deps --- Makefile | 9 ++-- tests/test_litellm/test_gate_slot_lock.py | 52 ++++++++++++++++------- 2 files changed, 43 insertions(+), 18 deletions(-) diff --git a/Makefile b/Makefile index 7fe5d1f8045..bdb643e3ec9 100644 --- a/Makefile +++ b/Makefile @@ -4,7 +4,7 @@ .PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \ test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \ test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \ - info lint lint-dev lint-checks format \ + info lint lint-inner lint-dev lint-checks format \ lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \ lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \ install-dev install-proxy-dev install-test-deps install-hooks \ @@ -239,8 +239,11 @@ check-import-safety: $(LINT_DEP_INSTALL) # does (merge-base with origin/litellm_internal_staging). Setup (env sync, Prisma client, # base fetch) runs once up front; the checks themselves are independent, so a sub-make # fans them out with -j and the fast ones finish under basedpyright's shadow. -lint: lint-install lint-fetch-base - $(GATE_SLOT_LOCK) $(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks +lint: + @$(GATE_SLOT_LOCK) $(MAKE) lint-inner + +lint-inner: lint-install lint-fetch-base + $(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_E2E_DEP_INSTALL= LINT_DEP_BASE= lint-checks lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright lint-e2e-basedpyright check-circular-imports check-import-safety diff --git a/tests/test_litellm/test_gate_slot_lock.py b/tests/test_litellm/test_gate_slot_lock.py index 1cf52ae89f6..17fa8547ce7 100644 --- a/tests/test_litellm/test_gate_slot_lock.py +++ b/tests/test_litellm/test_gate_slot_lock.py @@ -115,12 +115,8 @@ def test_two_slots_admit_two_holders_at_once(tmp_path: Path) -> None: first_started = tmp_path / "first.started" second_started = tmp_path / "second.started" env = _env(lock_dir, "2") - first = subprocess.Popen( - _wrapped([START_THEN_WAIT_FOR, str(first_started), str(second_started)]), env=env - ) - second = subprocess.Popen( - _wrapped([START_THEN_WAIT_FOR, str(second_started), str(first_started)]), env=env - ) + first = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(first_started), str(second_started)]), env=env) + second = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(second_started), str(first_started)]), env=env) assert first.wait(timeout=30) == 0 assert second.wait(timeout=30) == 0 @@ -131,9 +127,7 @@ def test_contender_beyond_capacity_queues_until_the_slot_frees(tmp_path: Path) - release = tmp_path / "release" done = tmp_path / "done" env = _env(lock_dir, "1") - holder = subprocess.Popen( - _wrapped([START_THEN_WAIT_FOR, str(holder_started), str(release)]), env=env - ) + holder = subprocess.Popen(_wrapped([START_THEN_WAIT_FOR, str(holder_started), str(release)]), env=env) try: assert _wait_until(holder_started.exists, 10) contender = subprocess.Popen( @@ -274,9 +268,7 @@ def test_killed_holder_releases_its_slot_for_the_next_contender(tmp_path: Path) assert b"freed" in after.stdout -def test_acquire_slot_holds_marks_and_releases_in_process( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: +def test_acquire_slot_holds_marks_and_releases_in_process(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: lock_dir = tmp_path / "locks" monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "") monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir)) @@ -293,9 +285,7 @@ def test_acquire_slot_holds_marks_and_releases_in_process( fcntl.flock(probe, fcntl.LOCK_UN) -def test_held_slot_context_manager_releases_on_exit( - tmp_path: Path, monkeypatch: pytest.MonkeyPatch -) -> None: +def test_held_slot_context_manager_releases_on_exit(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: lock_dir = tmp_path / "locks" monkeypatch.setenv("LITELLM_GATE_SLOT_HELD", "") monkeypatch.setenv("LITELLM_GATE_SLOT_DIR", str(lock_dir)) @@ -309,3 +299,35 @@ def test_held_slot_context_manager_releases_on_exit( with (lock_dir / "slot-0.lock").open("wb") as probe: fcntl.flock(probe, fcntl.LOCK_EX | fcntl.LOCK_NB) fcntl.flock(probe, fcntl.LOCK_UN) + + +def _make_rule(target: str) -> tuple[list[str], list[str]]: + database = subprocess.run( + ["make", "--dry-run", "--print-data-base", "info"], + cwd=ROOT, + capture_output=True, + text=True, + check=True, + ).stdout + lines = database.splitlines() + for index, line in enumerate(lines): + if line != f"{target}:" and not line.startswith(f"{target}: "): + continue + recipe: list[str] = [] + for follower in lines[index + 1 :]: + if follower.startswith("#"): + continue + if not follower.startswith("\t"): + break + recipe.append(follower.strip()) + return line.split(":", 1)[1].split(), recipe + raise AssertionError(f"target {target} not found in make database") + + +def test_direct_make_lint_takes_a_slot_before_any_setup() -> None: + lint_prerequisites, lint_recipe = _make_rule("lint") + assert lint_prerequisites == [] + assert any("$(GATE_SLOT_LOCK)" in line for line in lint_recipe) + inner_prerequisites, _ = _make_rule("lint-inner") + assert "lint-install" in inner_prerequisites + assert "lint-fetch-base" in inner_prerequisites From b84dd6922e8bb8848a0ad6704ceba856a944b6a7 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Tue, 11 Aug 2026 20:10:46 -0400 Subject: [PATCH 08/19] fix(bedrock): stop leaking managed-batch litellm_params to the provider A Bedrock managed-batch deployment carries aws_batch_role_arn, s3_bucket_name, s3_region_name, s3_output_bucket_name and bedrock_tags in its litellm_params, and the batch and files transformations read all five from there. None was registered in all_litellm_params, so the param builder swept them into extra_body on every other route that deployment serves: Bedrock answers "aws_batch_role_arn: Extra inputs are not permitted" on Anthropic models and "extraneous key [aws_batch_role_arn] is not permitted" on Nova, Llama and Titan, so configuring batch turns every chat and embedding request to that model into a 400. Register them alongside the agentic-loop and callback-credential fields, which are listed for exactly this reason. The batch path is unaffected because GenericLiteLLMParams is extra="allow" and preserves them into litellm_params for the transformations that consume them. Before this, batch could only be configured on a deployment dedicated to batch; the same model group could not serve both. --- litellm/types/utils.py | 13 ++++++++++++ tests/test_litellm/test_utils.py | 36 +++++++++++++++++++++++++++++++- 2 files changed, 48 insertions(+), 1 deletion(-) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 9baadc36f6b..b165aceb269 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3398,9 +3398,22 @@ agentic_loop_internal_litellm_params: Final = [ # the provider. TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" +# Bedrock managed-batch deployment config, read from litellm_params by the batch and +# files transformations. Listed for the same reason as the fields above: these sit on +# a deployment that also serves chat, so leaking them into extra_body makes Bedrock +# reject every non-batch request to that deployment. +bedrock_batch_litellm_params: Final = [ + "aws_batch_role_arn", + "s3_bucket_name", + "s3_region_name", + "s3_output_bucket_name", + "bedrock_tags", +] + all_litellm_params = ( agentic_loop_internal_litellm_params + [TRUSTED_CALLBACK_VARS_FIELD] + + bedrock_batch_litellm_params + [ "metadata", "litellm_metadata", diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 01e7e5c7ffd..27e7096189b 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -29,7 +29,8 @@ from litellm.types.utils import ( StreamingChoices, Usage, ) -from litellm.types.utils import all_litellm_params +from litellm.types.utils import all_litellm_params, bedrock_batch_litellm_params +from litellm.types.router import GenericLiteLLMParams from litellm.utils import ( ProviderConfigManager, TextCompletionStreamWrapper, @@ -4752,3 +4753,36 @@ def test_websearch_interception_control_fields_never_reach_the_provider(): f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}" ) assert set(WEBSEARCH_INTERNAL_CONTROL_FIELDS) <= set(all_litellm_params) + + +def test_bedrock_batch_params_never_reach_the_provider(): + """A Bedrock managed-batch deployment carries aws_batch_role_arn / s3_* / + bedrock_tags in its litellm_params, and the same deployment also serves chat. + Anything the param builder does not recognize is swept into extra_body, so + Bedrock rejects the whole call: `aws_batch_role_arn: Extra inputs are not + permitted` (Anthropic models) or `extraneous key [aws_batch_role_arn] is not + permitted` (Nova/Llama/Titan), turning every non-batch request to that + deployment into a 400. + + The batch path is unaffected by registering them, because GenericLiteLLMParams + is extra="allow" and preserves them into litellm_params for the batch and files + transformations that read them. + """ + kwargs = { + "a_real_provider_specific_param": 1, + **{field: "configured-value" for field in bedrock_batch_litellm_params}, + } + + non_default = get_non_default_completion_params(dict(kwargs)) + + assert non_default == {"a_real_provider_specific_param": 1}, ( + "bedrock batch params leaked into the provider params: " + f"{sorted(set(non_default) - {'a_real_provider_specific_param'})}" + ) + assert set(bedrock_batch_litellm_params) <= set(all_litellm_params) + + batch_params = dict(GenericLiteLLMParams(**kwargs)) + assert all(batch_params.get(field) == "configured-value" for field in bedrock_batch_litellm_params), ( + "registering these must not strip them from the batch path: " + f"{sorted(f for f in bedrock_batch_litellm_params if batch_params.get(f) != 'configured-value')}" + ) From 0c5c9c79d746449be950313d860c9704c6aa57f4 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Wed, 12 Aug 2026 03:39:46 -0400 Subject: [PATCH 09/19] fix(bedrock): carry s3_output_bucket_name and bedrock_tags through credential normalization Registering the five managed-batch fields in all_litellm_params stops them leaking into extra_body, but two of them never reached the transformation that reads them. CredentialLiteLLMParams is a whitelist, so get_deployment_credentials_with_provider round-tripped the deployment and silently dropped s3_output_bucket_name and bedrock_tags before the files/batch/passthrough callers saw them. s3_bucket_name, s3_region_name and aws_batch_role_arn were added to that model for #25104; these two are the remainder of the same deployment config bedrock_tags is typed as a plain list rather than a stricter shape so a malformed value still reaches _validate_bedrock_tags and gets its own error message instead of a Pydantic one The preservation assertion previously round-tripped through GenericLiteLLMParams, which is extra="allow" and would hold even for a field nothing declares. It now also reproduces the CredentialLiteLLMParams normalization the proxy actually performs, and fails naming exactly the dropped fields without this change --- litellm/types/router.py | 11 +++++++++++ tests/test_litellm/test_utils.py | 29 +++++++++++++++++++++++------ 2 files changed, 34 insertions(+), 6 deletions(-) diff --git a/litellm/types/router.py b/litellm/types/router.py index 217364c48b7..e0ab40b6973 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -266,6 +266,17 @@ class CredentialLiteLLMParams(BaseModel): s3_region_name: str | None = None s3_encryption_key_id: str | None = None aws_batch_role_arn: str | None = None + # Same reason as ``azure_ad_token`` above: this model is a whitelist, so a + # managed-batch field it does not declare is silently dropped by + # ``get_deployment_credentials_with_provider`` before the batch and files + # transformations that read it ever run. ``s3_bucket_name`` / + # ``s3_region_name`` / ``aws_batch_role_arn`` were added for #25104; these two + # are the remainder of the same deployment config. + s3_output_bucket_name: str | None = None + # A list of {"key": str, "value": str}; the batch transformation validates the + # shape itself via _validate_bedrock_tags, so this stays a plain list to keep + # that error message rather than failing earlier with a Pydantic one. + bedrock_tags: list | None = None ## IBM WATSONX ## watsonx_region_name: str | None = None diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 27e7096189b..a8640e252bb 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -30,7 +30,7 @@ from litellm.types.utils import ( Usage, ) from litellm.types.utils import all_litellm_params, bedrock_batch_litellm_params -from litellm.types.router import GenericLiteLLMParams +from litellm.types.router import CredentialLiteLLMParams, GenericLiteLLMParams from litellm.utils import ( ProviderConfigManager, TextCompletionStreamWrapper, @@ -4768,10 +4768,13 @@ def test_bedrock_batch_params_never_reach_the_provider(): is extra="allow" and preserves them into litellm_params for the batch and files transformations that read them. """ - kwargs = { - "a_real_provider_specific_param": 1, - **{field: "configured-value" for field in bedrock_batch_litellm_params}, + # bedrock_tags is a list of {"key", "value"} dicts; the rest are plain strings, so + # give each field a value of its real shape rather than one string for all of them. + configured = { + field: ([{"key": "team", "value": "configured-value"}] if field == "bedrock_tags" else "configured-value") + for field in bedrock_batch_litellm_params } + kwargs = {"a_real_provider_specific_param": 1, **configured} non_default = get_non_default_completion_params(dict(kwargs)) @@ -4782,7 +4785,21 @@ def test_bedrock_batch_params_never_reach_the_provider(): assert set(bedrock_batch_litellm_params) <= set(all_litellm_params) batch_params = dict(GenericLiteLLMParams(**kwargs)) - assert all(batch_params.get(field) == "configured-value" for field in bedrock_batch_litellm_params), ( + assert all(batch_params.get(field) == configured[field] for field in bedrock_batch_litellm_params), ( "registering these must not strip them from the batch path: " - f"{sorted(f for f in bedrock_batch_litellm_params if batch_params.get(f) != 'configured-value')}" + f"{sorted(f for f in bedrock_batch_litellm_params if batch_params.get(f) != configured[f])}" + ) + + # GenericLiteLLMParams is extra="allow", so the assertion above would hold even for a + # field nothing declares. The proxy's files/batch/passthrough callers do not see that + # dict: get_deployment_credentials_with_provider round-trips the deployment through + # CredentialLiteLLMParams, which is a whitelist, so an undeclared field is dropped + # before the batch transformation reads it. Reproduce that round-trip here so the + # preservation claim covers the path the proxy actually takes. + normalized = CredentialLiteLLMParams.model_validate( + GenericLiteLLMParams(**kwargs).model_dump(exclude_none=True) + ).model_dump(exclude_none=True) + assert all(normalized.get(field) == configured[field] for field in bedrock_batch_litellm_params), ( + "credential normalization dropped batch params before the transformation: " + f"{sorted(f for f in bedrock_batch_litellm_params if normalized.get(f) != configured[f])}" ) From 2e33eab2b076426700b8815bca31e845af7925ee Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Fri, 14 Aug 2026 01:36:38 -0400 Subject: [PATCH 10/19] chore(ui): regenerate dashboard api types for the new bedrock batch params --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 0c128d9815b..cf37709c377 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27010,6 +27010,8 @@ export interface components { aws_web_identity_token?: string | null; /** Azure Ad Token */ azure_ad_token?: string | null; + /** Bedrock Tags */ + bedrock_tags?: unknown[] | null; /** Budget Duration */ budget_duration?: string | null; /** Cache Creation Input Audio Token Cost */ @@ -27235,6 +27237,8 @@ export interface components { s3_bucket_name?: string | null; /** S3 Encryption Key Id */ s3_encryption_key_id?: string | null; + /** S3 Output Bucket Name */ + s3_output_bucket_name?: string | null; /** S3 Region Name */ s3_region_name?: string | null; /** Search Context Cost Per Query */ @@ -35919,6 +35923,8 @@ export interface components { aws_web_identity_token?: string | null; /** Azure Ad Token */ azure_ad_token?: string | null; + /** Bedrock Tags */ + bedrock_tags?: unknown[] | null; /** Budget Duration */ budget_duration?: string | null; /** Cache Creation Input Audio Token Cost */ @@ -36144,6 +36150,8 @@ export interface components { s3_bucket_name?: string | null; /** S3 Encryption Key Id */ s3_encryption_key_id?: string | null; + /** S3 Output Bucket Name */ + s3_output_bucket_name?: string | null; /** S3 Region Name */ s3_region_name?: string | null; /** Search Context Cost Per Query */ From 2a75381a9fc0f9aac2c03274e1fdfee68f70aa07 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 11:44:06 -0700 Subject: [PATCH 11/19] style(batches): sort the common_utils import block --- litellm/proxy/batches_endpoints/endpoints.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 9a6bb054d1f..87f9927b191 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -23,9 +23,9 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, + add_internal_model_credentials, apply_team_provider_credentials, decode_model_from_file_id, - add_internal_model_credentials, encode_batch_response_ids, encode_file_id_with_model, ensure_batch_response_managed_file_ids, From e46ff74bc00f1249cd56ad982272f4d1491b21b1 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 11:46:01 -0700 Subject: [PATCH 12/19] fix(types): make bedrock_batch_litellm_params a tuple to satisfy the LIT002 lint gate --- litellm/types/utils.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index b165aceb269..272fbabf807 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3402,18 +3402,17 @@ TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars" # files transformations. Listed for the same reason as the fields above: these sit on # a deployment that also serves chat, so leaking them into extra_body makes Bedrock # reject every non-batch request to that deployment. -bedrock_batch_litellm_params: Final = [ +bedrock_batch_litellm_params: Final = ( "aws_batch_role_arn", "s3_bucket_name", "s3_region_name", "s3_output_bucket_name", "bedrock_tags", -] +) all_litellm_params = ( agentic_loop_internal_litellm_params - + [TRUSTED_CALLBACK_VARS_FIELD] - + bedrock_batch_litellm_params + + [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params] + [ "metadata", "litellm_metadata", From 71d951bfc0e9bbeb552009e9c1a9756a62100a71 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 11:56:04 -0700 Subject: [PATCH 13/19] chore(types): drop redundant comments around the bedrock batch params --- litellm/types/router.py | 9 --------- tests/test_litellm/test_utils.py | 8 -------- 2 files changed, 17 deletions(-) diff --git a/litellm/types/router.py b/litellm/types/router.py index e0ab40b6973..f3f9276e6ba 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -266,16 +266,7 @@ class CredentialLiteLLMParams(BaseModel): s3_region_name: str | None = None s3_encryption_key_id: str | None = None aws_batch_role_arn: str | None = None - # Same reason as ``azure_ad_token`` above: this model is a whitelist, so a - # managed-batch field it does not declare is silently dropped by - # ``get_deployment_credentials_with_provider`` before the batch and files - # transformations that read it ever run. ``s3_bucket_name`` / - # ``s3_region_name`` / ``aws_batch_role_arn`` were added for #25104; these two - # are the remainder of the same deployment config. s3_output_bucket_name: str | None = None - # A list of {"key": str, "value": str}; the batch transformation validates the - # shape itself via _validate_bedrock_tags, so this stays a plain list to keep - # that error message rather than failing earlier with a Pydantic one. bedrock_tags: list | None = None ## IBM WATSONX ## watsonx_region_name: str | None = None diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index a8640e252bb..661b6ed7244 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -4768,8 +4768,6 @@ def test_bedrock_batch_params_never_reach_the_provider(): is extra="allow" and preserves them into litellm_params for the batch and files transformations that read them. """ - # bedrock_tags is a list of {"key", "value"} dicts; the rest are plain strings, so - # give each field a value of its real shape rather than one string for all of them. configured = { field: ([{"key": "team", "value": "configured-value"}] if field == "bedrock_tags" else "configured-value") for field in bedrock_batch_litellm_params @@ -4790,12 +4788,6 @@ def test_bedrock_batch_params_never_reach_the_provider(): f"{sorted(f for f in bedrock_batch_litellm_params if batch_params.get(f) != configured[f])}" ) - # GenericLiteLLMParams is extra="allow", so the assertion above would hold even for a - # field nothing declares. The proxy's files/batch/passthrough callers do not see that - # dict: get_deployment_credentials_with_provider round-trips the deployment through - # CredentialLiteLLMParams, which is a whitelist, so an undeclared field is dropped - # before the batch transformation reads it. Reproduce that round-trip here so the - # preservation claim covers the path the proxy actually takes. normalized = CredentialLiteLLMParams.model_validate( GenericLiteLLMParams(**kwargs).model_dump(exclude_none=True) ).model_dump(exclude_none=True) From 1524880dceefe90d0ba92f03710a24ae609d931b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 12:09:59 -0700 Subject: [PATCH 14/19] fix(batches): sign the retrieve-path output read with the deployment's AWS credentials --- litellm/batches/batch_utils.py | 6 +- .../test_litellm/batches/test_batch_utils.py | 56 +++++++++++++++++++ 2 files changed, 60 insertions(+), 2 deletions(-) diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index f7aa6c50de8..ebef60b41c9 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -5,6 +5,7 @@ from typing import Any, Final, Literal import litellm from litellm._logging import verbose_logger +from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS from litellm.litellm_core_utils.llm_cost_calc.utils import _parse_prompt_tokens_details from litellm.types.llms.openai import Batch from litellm.types.utils import CallTypes, ModelInfo, Usage @@ -295,7 +296,7 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: if litellm_params: # List of credential keys that should be passed to file operations - credential_keys: Final = [ + credential_keys: Final = ( "api_key", "api_base", "api_version", @@ -310,7 +311,8 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: "timeout", "max_retries", "_litellm_internal_model_credentials", - ] + *AWS_CREDENTIAL_KWARGS_KEYS, + ) for key in credential_keys: if key in litellm_params: credentials[key] = litellm_params[key] diff --git a/tests/test_litellm/batches/test_batch_utils.py b/tests/test_litellm/batches/test_batch_utils.py index cacae3624f3..d2074853f2b 100644 --- a/tests/test_litellm/batches/test_batch_utils.py +++ b/tests/test_litellm/batches/test_batch_utils.py @@ -1243,3 +1243,59 @@ def test_extract_credentials_forwards_the_trusted_model_credential_snapshot(): credentials = bu._extract_file_access_credentials({"_litellm_internal_model_credentials": snapshot}) assert credentials["_litellm_internal_model_credentials"] is snapshot + + +def test_extract_credentials_forwards_the_deployment_aws_credentials(): + """The retrieve path's logging object carries the deployment's AWS keys in its + litellm_params, and the S3 read of the output file signs with whatever afile_content + receives. Dropping them here sent the read to the ambient credential chain, so a + deployment whose only AWS credentials live in its litellm_params never recorded + batch cost on retrieve even once the bucket resolved.""" + params = { + "aws_access_key_id": "AKIA-deployment", + "aws_secret_access_key": "secret-deployment", + "aws_session_token": "token-deployment", + "aws_region_name": "us-west-2", + "aws_role_name": "arn:aws:iam::123456789012:role/batch-reader", + "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + } + + credentials = bu._extract_file_access_credentials(params) + + assert credentials == {key: value for key, value in params.items() if key != "model"} + + +@pytest.mark.asyncio +async def test_output_file_content_bedrock_reads_with_deployment_aws_credentials(monkeypatch): + import litellm.files.main as files_main + + captured: dict = {} + + async def fake_afile_content(**kw): + captured.update(kw) + return type("R", (), {"content": b""})() + + monkeypatch.setattr(files_main, "afile_content", fake_afile_content) + snapshot = MappingProxyType({"s3_bucket_name": "configured-bucket", "aws_region_name": "us-west-2"}) + + await bu._fetch_batch_output_file_content( + _batch("s3://configured-bucket/litellm-batch-outputs/job-1/out.jsonl.out"), + custom_llm_provider="bedrock", + litellm_params={ + "aws_access_key_id": "AKIA-deployment", + "aws_secret_access_key": "secret-deployment", + "aws_session_token": "token-deployment", + "aws_region_name": "us-west-2", + "_litellm_internal_model_credentials": snapshot, + "model": "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + }, + ) + + assert captured["file_id"] == "s3://configured-bucket/litellm-batch-outputs/job-1/out.jsonl.out" + assert captured["custom_llm_provider"] == "bedrock" + assert captured["aws_access_key_id"] == "AKIA-deployment" + assert captured["aws_secret_access_key"] == "secret-deployment" + assert captured["aws_session_token"] == "token-deployment" + assert captured["aws_region_name"] == "us-west-2" + assert captured["_litellm_internal_model_credentials"] is snapshot + assert "model" not in captured From 4eadf92adee832aa1ef3e52af6b66787614eda0d Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 15 Aug 2026 14:56:34 -0700 Subject: [PATCH 15/19] feat(mcp): scope gateway session bearers to the RFC 8707 resource (#35045) --- .../mcp_server/auth/user_api_key_auth_mcp.py | 11 +- .../mcp_server/discoverable_endpoints.py | 4 + .../mcp_server/gateway_dcr_flow.py | 92 ++++++- .../mcp_server/mcp_server_manager.py | 22 +- .../_experimental/mcp_server/oauth_utils.py | 4 +- .../outbound_credentials/session_token.py | 16 +- litellm/proxy/_types.py | 8 + .../auth/test_user_api_key_auth_mcp.py | 65 ++++- .../mcp_server/test_gateway_dcr_flow.py | 244 ++++++++++++++++++ .../test_mcp_oauth_passthrough_cold_start.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 63 +++++ 11 files changed, 515 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 95a3806e8ad..d13b39661ad 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -83,7 +83,7 @@ class UnloadableEntitlementError(Exception): def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | None = None) -> list[str] | None: """Resolve the single MCP server name a cold-start passthrough bypass may target. Delegates parsing to - :meth:`MCPRequestHandler._extract_target_server_names_from_path` so the + :meth:`MCPRequestHandler.extract_target_server_names_from_path` so the names used here always match the names downstream routing uses; returns ``None`` whenever the bypass must not activate (aggregate ``/mcp``, multi-server CSV paths, or any other unrecognized path). @@ -94,7 +94,7 @@ def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: list[str] | header/path mismatch here is a sign of a confused or hostile caller — refuse the cold-start bypass rather than admit anonymously based on the path while the header advertises a stricter, non-passthrough target.""" - servers: Final = MCPRequestHandler._extract_target_server_names_from_path(path) + servers: Final = MCPRequestHandler.extract_target_server_names_from_path(path) if len(servers) != 1: verbose_logger.debug( "MCP cold-start: path %r resolved to %r; passthrough 401 bypass " @@ -215,7 +215,7 @@ def _is_gateway_dcr_challenge_scope( return False if _has_client_supplied_mcp_auth(mcp_auth_header, mcp_server_auth_headers): return False - if len(MCPRequestHandler._extract_target_server_names_from_path(route)) == 0: + if len(MCPRequestHandler.extract_target_server_names_from_path(route)) == 0: return True return _gateway_dcr_challenge_target(route, mcp_servers, client_ip) is not None @@ -579,7 +579,7 @@ class MCPRequestHandler: return oauth2_headers, raw_headers, mcp_auth_header, mcp_server_auth_headers @staticmethod - def _extract_target_server_names_from_path(path: str) -> list[str]: + def extract_target_server_names_from_path(path: str) -> list[str]: """ Extract the target MCP server name(s) from the standard MCP transport URL patterns: ``/mcp/{server_name_or_csv}[/...]`` and @@ -836,6 +836,7 @@ class MCPRequestHandler: case SessionBearerAdmitted(): try: admitted: Final = await MCPRequestHandler._reload_admitted_user(result.principal.user_id) + admitted.mcp_session_resource_server_id = result.principal.resource_server_id await MCPRequestHandler._enforce_admitted_live_policy( admitted=admitted, request=request, route=route ) @@ -1168,7 +1169,7 @@ class MCPRequestHandler: (header/path TOCTOU). For non-``/mcp/...`` paths (where the path does not encode targets), fall back to the header. """ - path_targets: Final = MCPRequestHandler._extract_target_server_names_from_path(path) + path_targets: Final = MCPRequestHandler.extract_target_server_names_from_path(path) if path_targets: return path_targets # Path did not resolve to /mcp/... targets — trust the header diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 693e3f8e47d..86e97b55a8e 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1655,6 +1655,7 @@ async def authorize( code_challenge_method: str | None = None, response_type: str | None = None, scope: str | None = None, + resource: str | None = None, ): # Redirect to real OAuth provider with PKCE support from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( @@ -1671,6 +1672,7 @@ async def authorize( code_challenge_method=code_challenge_method, response_type=response_type, session_user_id=_session_cookie_user_id(request), + resource=resource, ) lookup_name: Final[str | None] = mcp_server_name or client_id @@ -1721,6 +1723,7 @@ async def token_endpoint( code_verifier: str = Form(None), refresh_token: str | None = Form(None), scope: str | None = Form(None), + resource: str | None = Form(None), mcp_server_name: str | None = None, ): """ @@ -1753,6 +1756,7 @@ async def token_endpoint( master_key=master_key, reload_user=_reload_active_user_by_id, cache=user_api_key_cache, + resource=resource, ) lookup_name: Final = mcp_server_name or client_id diff --git a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py index 4c1b78c754a..85885fc75f5 100644 --- a/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py +++ b/litellm/proxy/_experimental/mcp_server/gateway_dcr_flow.py @@ -56,6 +56,8 @@ from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.proxy._experimental.mcp_server.oauth_utils import ( TOKEN_NO_CACHE_HEADERS, + canonical_resource_uri, + canonicalize_url_identity, get_request_base_url, is_loopback_redirect_host, validate_redirect_uri_shape, @@ -77,6 +79,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) +from litellm.types.mcp_server.mcp_server_manager import MCPServer GATEWAY_DCR_CLIENT_ID_PREFIX: Final = "llm_dcrc_" """Marker prefix on every gateway-issued DCR client_id so the root authorize/token @@ -169,6 +172,7 @@ class _ConnectFlow(BaseModel): code_challenge: str = Field(min_length=1) jti: str = Field(min_length=1) exp: int + resource_server_id: str | None = None class _GatewayAuthCode(BaseModel): @@ -185,6 +189,7 @@ class _GatewayAuthCode(BaseModel): jti: str = Field(min_length=1) iat: int exp: int + resource_server_id: str | None = None def is_gateway_dcr_client_id(client_id: str | None) -> bool: @@ -204,7 +209,13 @@ def _oauth_error(status_code: int, error: str, description: str) -> JSONResponse def _seal(prefix: str, payload: BaseModel) -> str: - return prefix + encrypt_value_helper(payload.model_dump_json()) + """Serialized ``exclude_none`` for the same reason session JWTs are minted that way: an + optional claim that is unset never reaches the wire, so during a rolling deploy a blob + sealed by a new pod without the new claim set stays byte-compatible with predating pods + whose strict models forbid unknown keys. This holds for every sealed artifact and every + future optional claim by construction; it requires each optional field to default to + ``None`` so reopening restores exactly what was sealed.""" + return prefix + encrypt_value_helper(payload.model_dump_json(exclude_none=True)) _SealedModelT = TypeVar("_SealedModelT", bound=BaseModel) @@ -320,6 +331,44 @@ def relative_request_url(request: Request) -> str: return f"{path}?{request.url.query}" if request.url.query else path +def resolve_scoped_resource_server(request: Request, resource: str | None) -> MCPServer | None: + """Resolve an RFC 8707 ``resource`` value to the single gateway-managed oauth2 server it + names, or ``None`` for every other shape: absent, the aggregate resource, a foreign + host, an unparseable value, a multi-server path, an unknown name, or any server mode the + keyless gateway flow does not serve (whose protected-resource metadata never directs a + client here). ``None`` means the flow stays unscoped and byte-identical to today, so a + hostile or confused ``resource`` can never widen anything; a resolved server only ever + NARROWS the session via the sealed scope. + + Resolution is an IDENTITY question, deliberately free of the per-IP visibility filter: + access is enforced where it belongs (grant intersection at admission, IP checks on the + MCP routes), while filtering here would mint an entitlement-wide UNSCOPED bearer exactly + when the caller asked to narrow, and would let authorize-time vs token-time IP drift + turn a matching redemption into a spurious ``invalid_target``.""" + if resource is None: + return None + canonical: Final = canonical_resource_uri(resource) + if canonical is None: + return None + base: Final = canonicalize_url_identity(get_request_base_url(request)) + if canonical == f"{base}/mcp" or not canonical.startswith(f"{base}/"): + return None + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( # noqa: PLC0415 # proxy import cycle + MCPRequestHandler, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # proxy import cycle + global_mcp_server_manager, + ) + + names: Final = MCPRequestHandler.extract_target_server_names_from_path(canonical[len(base) :]) + if len(names) != 1: + return None + server: Final = global_mcp_server_manager.get_mcp_server_by_name(names[0]) + if server is None or not server.is_gateway_managed_oauth2: + return None + return server + + def aggregate_authorize( request: Request, client_id: str, @@ -329,11 +378,16 @@ def aggregate_authorize( code_challenge_method: str | None, response_type: str | None, session_user_id: str | None, + resource: str | None = None, ) -> Response: """The aggregate authorize verb: validate the client, require S256 PKCE, interpose LiteLLM sign-in, and hand the browser to the connect page with the flow sealed into a per-flow cookie. + A per-server RFC 8707 ``resource`` naming a gateway-managed oauth2 server scopes the + flow to that one server: the scope is sealed into the flow, carried into the code, and + bound into the session token, while the connect page interlude runs exactly as before. + Validation failures respond directly with 400 and never redirect: per RFC 6749 section 4.1.2.1 an unvalidated redirect URI must not receive an error redirect, and once the client is at fault there is no trusted place to send the browser. @@ -358,6 +412,7 @@ def aggregate_authorize( login_url: Final = f"{base_url}/sso/key/generate?{urlencode({'return_to': relative_request_url(request)})}" return RedirectResponse(login_url, status_code=303) now: Final = datetime.now(timezone.utc) + scoped_server: Final = resolve_scoped_resource_server(request, resource) handle: Final = secrets.token_urlsafe(24) flow: Final = _ConnectFlow( user_id=session_user_id, @@ -367,6 +422,7 @@ def aggregate_authorize( code_challenge=code_challenge, jti=secrets.token_urlsafe(24), exp=int(now.timestamp()) + CONNECT_FLOW_TTL_SECONDS, + resource_server_id=scoped_server.server_id if scoped_server is not None else None, ) connect_url: Final = _append_query_params( f"{base_url}/ui/connect", @@ -455,6 +511,7 @@ async def complete_connect_flow( jti=secrets.token_urlsafe(24), iat=int(now.timestamp()), exp=int(now.timestamp()) + code_ttl, + resource_server_id=flow.resource_server_id, ), ) params: Final = {"code": code, **({"state": flow.state} if flow.state else {})} @@ -587,6 +644,20 @@ def _reload_failure_response(failure: ReloadUserFailure) -> Response: assert_never(failure) +def _resource_conflicts_with_scope( + request: Request, resource: str | None, sealed_resource_server_id: str | None +) -> bool: + """True when a scoped grant is being redeemed for a DIFFERENT resource than the one + sealed into it (RFC 8707 section 2.2: reject with ``invalid_target``). An absent + ``resource`` never conflicts (the sealed scope still binds the minted session), and an + unscoped grant ignores the parameter entirely, exactly as the endpoint always has, so + no pre-existing client breaks.""" + if sealed_resource_server_id is None or resource is None: + return False + resolved: Final = resolve_scoped_resource_server(request, resource) + return resolved is None or resolved.server_id != sealed_resource_server_id + + async def aggregate_token( request: Request, grant_type: str, @@ -598,6 +669,7 @@ async def aggregate_token( master_key: str | None, reload_user: ReloadUser, cache: DualCache, + resource: str | None = None, ) -> Response: """The aggregate token verb: authorization_code and refresh_token grants for the identity-only session pair. Every path re-validates the litellm user live before @@ -609,10 +681,12 @@ async def aggregate_token( now: Final = datetime.now(timezone.utc) if grant_type == "authorization_code": return await _authorization_code_grant( + request=request, code=code, redirect_uri=redirect_uri, client_id=client_id, code_verifier=code_verifier, + resource=resource, keys=keys, now=now, reload_user=reload_user, @@ -620,8 +694,10 @@ async def aggregate_token( ) if grant_type == "refresh_token": return await _refresh_token_grant( + request=request, refresh_token=refresh_token, client_id=client_id, + resource=resource, keys=keys, now=now, reload_user=reload_user, @@ -631,10 +707,12 @@ async def aggregate_token( async def _authorization_code_grant( + request: Request, code: str | None, redirect_uri: str | None, client_id: str, code_verifier: str | None, + resource: str | None, keys: SessionKeys, now: datetime, reload_user: ReloadUser, @@ -651,6 +729,8 @@ async def _authorization_code_grant( return _oauth_error(400, "invalid_grant", "the authorization code has expired") if client_id != parsed.client_id or redirect_uri != parsed.redirect_uri: return _oauth_error(400, "invalid_grant", "the authorization code was issued to a different client") + if _resource_conflicts_with_scope(request, resource, parsed.resource_server_id): + return _oauth_error(400, "invalid_target", "resource does not match the scope this code was issued for") if not _pkce_verifier_matches(code_verifier, parsed.code_challenge): return _oauth_error(400, "invalid_grant", "PKCE verification failed") # Revalidate the user BEFORE claiming the code, so a transient DB outage (a retryable @@ -666,12 +746,18 @@ async def _authorization_code_grant( parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS, ): return _oauth_error(400, "invalid_grant", "the authorization code was already used") - return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now) + return _session_token_pair( + SessionPrincipal(user_id=parsed.user_id, client_id=client_id, resource_server_id=parsed.resource_server_id), + keys, + now, + ) async def _refresh_token_grant( + request: Request, refresh_token: str | None, client_id: str, + resource: str | None, keys: SessionKeys, now: datetime, reload_user: ReloadUser, @@ -682,6 +768,8 @@ async def _refresh_token_grant( opened: Final = open_session_refresh_bearer(refresh_token, keys, now, expected_client_id=client_id) if not isinstance(opened, SessionRefreshOpened): return _oauth_error(400, "invalid_grant", "the refresh token is invalid for this client") + if _resource_conflicts_with_scope(request, resource, opened.principal.resource_server_id): + return _oauth_error(400, "invalid_target", "resource does not match the scope this token was issued for") failure: Final = await reload_user(opened.principal.user_id) if failure is not None: return _reload_failure_response(failure) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4f94f94acb5..c782f0dfa09 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2491,6 +2491,18 @@ class MCPServerManager: open_ids.update(submitted_server_ids) return open_ids + @staticmethod + def _admitted_session_resource_scope(user_api_key_auth: UserAPIKeyAuth | None) -> str | None: + """The single server an admitted session subject's bearer was scoped to at authorize + time (RFC 8707 resource), or None for every other principal shape and for unscoped + sessions. Read at every return path of :meth:`get_allowed_mcp_servers`, including + the exception fallback, and applied AFTER every union (grants, operator-open, + submitted) because the scope is a ceiling over the whole reachable set; a resolver + fault therefore never widens a scoped bearer to the allow-all set.""" + if user_api_key_auth is None or not _is_mcp_admitted_user_subject(user_api_key_auth): + return None + return user_api_key_auth.mcp_session_resource_server_id + async def get_allowed_mcp_servers(self, user_api_key_auth: UserAPIKeyAuth | None = None) -> list[str]: """ Get the allowed MCP Servers for the user. @@ -2600,13 +2612,19 @@ class MCPServerManager: if len(combined_servers) == 0: verbose_logger.debug("No allowed MCP Servers found for user api key auth.") - return list(combined_servers) + scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth) + return [server_id for server_id in combined_servers if scope is None or server_id == scope] except Exception: # noqa: BLE001 verbose_logger.exception( "Failed to get allowed MCP servers; team-level object_permission " "grants may be dropped. Falling back to global and submitted servers." ) - return list(dict.fromkeys(allow_all_server_ids + submitted_server_ids)) + scope = MCPServerManager._admitted_session_resource_scope(user_api_key_auth) + return [ + server_id + for server_id in dict.fromkeys(allow_all_server_ids + submitted_server_ids) + if scope is None or server_id == scope + ] async def resolve_toolset_tool_permissions( self, diff --git a/litellm/proxy/_experimental/mcp_server/oauth_utils.py b/litellm/proxy/_experimental/mcp_server/oauth_utils.py index 84b40b72258..a30b5ee9e49 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_utils.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_utils.py @@ -633,7 +633,7 @@ def canonicalize_url_identity(url: str) -> str: return urlunparse((scheme, netloc, parsed.path.rstrip("/"), "", "", "")) -def _canonical_resource_uri(url: str) -> str | None: +def canonical_resource_uri(url: str) -> str | None: """Canonicalize an upstream MCP server URL into an RFC 8707 resource identifier. Keeps only the scheme, host, port and path, which is the shape the MCP authorization spec's @@ -693,7 +693,7 @@ def resolve_upstream_resource(mcp_server: "MCPServer") -> str | None: mcp_server.server_id, ) return None - canonical: Final = _canonical_resource_uri(mcp_server.url) + canonical: Final = canonical_resource_uri(mcp_server.url) if canonical is None: verbose_logger.warning( "MCP server %s sets upstream_resource=auto but its url is not an absolute URI, so no " diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py index 3ef7327cda3..15f5f82c4b6 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -85,11 +85,18 @@ class SessionPrincipal(BaseModel): enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless, gateway-sealed) DCR client identifier the token was issued to; the token endpoint requires it to match on the refresh grant. + + ``resource_server_id`` is the single MCP server this session was authorized for when + the client requested a per-server RFC 8707 resource at authorize time, or ``None`` for + the aggregate scope. It is a RESTRICTION carried for admission to intersect against + the live grant resolution, never a grant by itself; the refresh grant re-mints from + this principal so the restriction survives rotation. """ model_config = ConfigDict(frozen=True) user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) + resource_server_id: str | None = None class SessionKeys(BaseModel): @@ -186,6 +193,7 @@ class _SessionClaims(BaseModel): kind: SessionTokenKind user_id: str = Field(min_length=1) client_id: str = Field(min_length=1) + resource_server_id: str | None = None def is_session_token(candidate: str) -> bool: @@ -286,9 +294,10 @@ def _mint( kind=kind, user_id=principal.user_id, client_id=principal.client_id, + resource_server_id=principal.resource_server_id, ) token: Final = prefix + jwt.encode( - claims.model_dump(), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM + claims.model_dump(exclude_none=True), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM ) size_bytes: Final = len(token.encode("utf-8")) if size_bytes > MAX_SESSION_TOKEN_BYTES: @@ -323,7 +332,10 @@ def _open( if now.timestamp() >= claims.exp: return SessionExpired() return OpenedSessionToken( - principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id), jti=claims.jti + principal=SessionPrincipal( + user_id=claims.user_id, client_id=claims.client_id, resource_server_id=claims.resource_server_id + ), + jti=claims.jti, ) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 934ac9ac3d7..a566d491597 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -2762,6 +2762,13 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # key off. Server-only and stripped from validated input for the same reason as the marker # above: a forged entry would let a caller pick which team's rpm bucket it is charged against. mcp_source_team_rpm_limits: dict[str, dict[str, int]] | None = Field(default=None, exclude=True) + # The single MCP server_id a gateway session bearer was scoped to at authorize time (RFC 8707 + # resource), or None for an aggregate-scope session. A RESTRICTION intersected against the live + # grant resolution, never a grant. Server-only, set exclusively by the MCP gateway admission + # path via post-construction assignment and stripped from validated input like the markers + # above; a forged value could at most narrow, but the stripping keeps the field's provenance + # single-owner so its meaning stays trustworthy. + mcp_session_resource_server_id: str | None = Field(default=None, exclude=True) via_virtual_key: bool = Field( default=False, exclude=True, @@ -2798,6 +2805,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # kwargs, model_validate, a JWT/key claim splat) so it can never be forged from caller data. values.pop("mcp_admitted_user_subject", None) values.pop("mcp_source_team_rpm_limits", None) + values.pop("mcp_session_resource_server_id", None) values.pop("via_virtual_key", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 52dc91ce24d..0209abee510 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -2646,7 +2646,7 @@ class TestMCPDelegateAuthToUpstream: def test_extract_target_server_names_matches_routing_parser(self): """ - Regression: _extract_target_server_names_from_path must match the + Regression: extract_target_server_names_from_path must match the downstream regex parser in server.py::_get_mcp_servers_in_path. Previously, a request to ``/mcp//garbage`` was parsed as @@ -2682,7 +2682,7 @@ class TestMCPDelegateAuthToUpstream: ("/", []), ] for path_input, expected in cases: - assert MCPRequestHandler._extract_target_server_names_from_path(path_input) == expected, ( + assert MCPRequestHandler.extract_target_server_names_from_path(path_input) == expected, ( f"path={path_input!r} → expected {expected!r}" ) assert (_get_mcp_servers_in_path(path_input) or []) == expected, ( @@ -8365,3 +8365,64 @@ class TestEntitlementFaultSemantics: ): allowed = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(allowed) == {"srv1"} + + +@pytest.mark.asyncio +class TestScopedSessionAdmission: + """LIT-4917: a session bearer sealed to one server (RFC 8707 resource at authorize) + carries that scope onto the admitted auth object, where the grant resolution intersects + it fail closed; an unscoped bearer carries None and is byte-identical to before.""" + + _MASTER_KEY = "sk-scoped-session-admission-master-key" + + def _bearer(self, resource_server_id): + from datetime import datetime, timezone + + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + session_keys_from_master_key, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SessionPrincipal, + mint_session_token, + ) + + keys = session_keys_from_master_key(self._MASTER_KEY) + principal = SessionPrincipal( + user_id="scoped-user", client_id="llm_dcrc_abc", resource_server_id=resource_server_id + ) + return mint_session_token(principal, keys, datetime(2030, 1, 1, tzinfo=timezone.utc)).token.get_secret_value() + + @pytest.mark.parametrize("scope", ["github-server-id", None]) + async def test_admission_carries_sealed_resource_scope(self, scope): + token = self._bearer(scope) + scope_dict = { + "type": "http", + "method": "POST", + "path": "/mcp/github", + "headers": [(b"host", b"testserver"), (b"authorization", f"Bearer {token}".encode())], + } + get_user_object = AsyncMock( + return_value=MagicMock( + user_id="scoped-user", + organization_id=None, + metadata={"scim_active": True}, + user_role=None, + object_permission=None, + object_permission_id=None, + tpm_limit=None, + rpm_limit=None, + ) + ) + with ( + patch("litellm.proxy.proxy_server.master_key", self._MASTER_KEY), + patch("litellm.proxy.auth.auth_checks.get_user_object", get_user_object), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + ): + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope_dict) + assert auth_result.mcp_admitted_user_subject is True + assert auth_result.mcp_session_resource_server_id == scope + + def test_scope_field_cannot_be_forged_through_construction(self): + forged = UserAPIKeyAuth(user_id="u1", mcp_session_resource_server_id="any-server") + assert forged.mcp_session_resource_server_id is None diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py index d8eab3ecb3b..cc65970a180 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_gateway_dcr_flow.py @@ -795,3 +795,247 @@ async def test_manual_delivery_page_renders_the_url_as_data_never_as_a_shell_com assert 'curl "' not in body assert "curl '" not in body assert 'value="' in body + + +def _scoped_mcp_server(name="github", **kw): + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + return MCPServer( + server_id=f"{name}-id", + name=name, + server_name=name, + alias=name, + url="https://upstream.example/mcp", + transport="http", + auth_type=MCPAuth.oauth2, + **kw, + ) + + +SCOPED_RESOURCE = "https://llm.example.com/mcp/github" +_MANAGER_PATCH = "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + + +def _scoped_authorize(client_id, resource, session_user_id="u1"): + return aggregate_authorize( + request=_request(query=f"client_id={client_id}"), + client_id=client_id, + redirect_uri=REDIRECT_URI, + state="client-state-123", + code_challenge=CODE_CHALLENGE, + code_challenge_method="S256", + response_type="code", + session_user_id=session_user_id, + resource=resource, + ) + + +async def _redeem(code, client_id, cache=None, **overrides): + arguments = { + "request": _request("/token", method="POST"), + "grant_type": "authorization_code", + "code": code, + "redirect_uri": REDIRECT_URI, + "client_id": client_id, + "code_verifier": CODE_VERIFIER, + "refresh_token": None, + "master_key": MASTER_KEY, + "reload_user": _reload_user_active, + "cache": cache or DualCache(), + } + return await aggregate_token(**{**arguments, **overrides}) + + +def _opened_principal(payload): + keys = session_keys_from_master_key(MASTER_KEY) + admitted = resolve_session_bearer(f"Bearer {payload['access_token']}", keys, datetime.now(timezone.utc)) + assert isinstance(admitted, SessionBearerAdmitted) + return admitted.principal + + +async def _finish_connect_page(response): + handle, cookies = _flow_cookie_from(response) + completed = await complete_connect_flow( + request=_request("/authorize/complete", cookies=cookies, method="POST"), + flow_handle=handle, + session_user_id="u1", + cache=DualCache(), + ) + return parse_qs(urlparse(completed.headers["location"]).query)["code"][0] + + +def _sealed_wire_json(sealed, prefix, debug_key): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + + raw = decrypt_value_helper(sealed.removeprefix(prefix), debug_key, return_original_value=False) + assert isinstance(raw, str) + return json.loads(raw) + + +@pytest.mark.asyncio +async def test_scoped_authorize_runs_connect_page_with_sealed_scope(): + """LIT-4917: a per-server RFC 8707 resource naming a gateway-managed oauth2 server + seals that server into the flow. The connect page interlude runs exactly as before + (the scope restricts, it never skips consent), and the code minted at the finish step + and the session pair it redeems for are both scoped.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server() + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + assert response.status_code == 303 + assert "/ui/connect" in response.headers["location"] + _, cookies = _flow_cookie_from(response) + assert _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow")["resource_server_id"] == "github-id" + code = await _finish_connect_page(response) + assert _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code")["resource_server_id"] == "github-id" + token_response = await _redeem(code, client_id) + assert token_response.status_code == 200 + principal = _opened_principal(json.loads(token_response.body)) + assert principal.resource_server_id == "github-id" + assert principal.user_id == "u1" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "resource, resolves", + [ + (None, False), + ("https://llm.example.com/mcp", False), + ("https://other.example.com/mcp/github", False), + ("https://llm.example.com/mcp/github,linear", False), + ("https://llm.example.com/mcp/unknown", None), + ("not a url", False), + ], +) +async def test_unscoped_resources_leave_flow_and_token_byte_identical(resource, resolves): + """Every resource shape outside 'exactly one gateway-managed server' keeps today's flow: + connect page interlude, and NONE of the minted artifacts carry the scope key on the + wire, not the flow cookie, not the code, not the session JWT, so an unscoped flow + started on a new pod completes on a pod whose strict models predate the claim.""" + import base64 + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = None if resolves is None else _scoped_mcp_server() + response = _scoped_authorize(client_id, resource) + assert response.status_code == 303 + assert "/ui/connect" in response.headers["location"] + _, cookies = _flow_cookie_from(response) + assert "resource_server_id" not in _sealed_wire_json(next(iter(cookies.values())), "", "gateway_connect_flow") + code = await _finish_connect_page(response) + assert "resource_server_id" not in _sealed_wire_json(code, GATEWAY_AUTH_CODE_PREFIX, "gateway_authorization_code") + token_response = await _redeem(code, client_id) + payload = json.loads(token_response.body) + assert _opened_principal(payload).resource_server_id is None + jwt_payload_segment = payload["access_token"].removeprefix("llm_session_").split(".")[1] + claims = json.loads(base64.urlsafe_b64decode(jwt_payload_segment + "=" * (-len(jwt_payload_segment) % 4))) + assert "resource_server_id" not in claims + + +@pytest.mark.asyncio +async def test_scoped_authorize_delegate_server_stays_unscoped(): + """A delegate-auth oauth2 server is outside the gateway-managed set (its keyless flow is + upstream PKCE via the relay), so a resource naming it never scopes the gateway flow.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server(delegate_auth_to_upstream=True) + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + assert "/ui/connect" in response.headers["location"] + code = await _finish_connect_page(response) + token_response = await _redeem(code, client_id) + assert _opened_principal(json.loads(token_response.body)).resource_server_id is None + + +@pytest.mark.asyncio +async def test_token_rejects_resource_conflicting_with_sealed_scope(): + """RFC 8707 section 2.2: redeeming a scoped code (or rotating a scoped refresh token) + for a DIFFERENT resource fails with invalid_target; an absent resource redeems fine and + the sealed scope still binds the minted pair, surviving refresh rotation.""" + from unittest.mock import patch + + client_id = (await _register([REDIRECT_URI]))["client_id"] + github = _scoped_mcp_server() + linear = _scoped_mcp_server(name="linear") + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = github + response = _scoped_authorize(client_id, SCOPED_RESOURCE) + code = await _finish_connect_page(response) + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = linear + mismatched = await _redeem(code, client_id, resource="https://llm.example.com/mcp/linear") + assert json.loads(mismatched.body)["error"] == "invalid_target" + + cache = DualCache() + token_response = await _redeem(code, client_id, cache=cache) + payload = json.loads(token_response.body) + assert _opened_principal(payload).resource_server_id == "github-id" + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = linear + refresh_mismatch = await _redeem( + None, + client_id, + cache=cache, + grant_type="refresh_token", + refresh_token=payload["refresh_token"], + resource="https://llm.example.com/mcp/linear", + ) + assert json.loads(refresh_mismatch.body)["error"] == "invalid_target" + + rotated = await _redeem( + None, client_id, cache=cache, grant_type="refresh_token", refresh_token=payload["refresh_token"] + ) + assert rotated.status_code == 200 + assert _opened_principal(json.loads(rotated.body)).resource_server_id == "github-id" + + +@pytest.mark.asyncio +async def test_resolve_scoped_resource_server_matrix(): + """Unit pin of the resource resolver: both per-server URL spellings resolve; the + aggregate resource, foreign hosts, CSV paths, unknown names, and non-gateway-managed + modes all return None so nothing outside the served set can enter the scoped flow.""" + from unittest.mock import patch + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server + + request = _request() + github = _scoped_mcp_server() + for resource, resolved_server, expected in [ + ("https://llm.example.com/mcp/github", github, "github-id"), + ("https://llm.example.com/github/mcp", github, "github-id"), + ("https://LLM.example.com/mcp/github/", github, "github-id"), + ("https://llm.example.com/mcp", github, None), + ("https://other.example.com/mcp/github", github, None), + ("https://llm.example.com/mcp/a,b", github, None), + ("https://llm.example.com/mcp/github", None, None), + ("https://llm.example.com/mcp/github", _scoped_mcp_server(delegate_auth_to_upstream=True), None), + (None, github, None), + ]: + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = resolved_server + result = resolve_scoped_resource_server(request, resource) + assert (result.server_id if result is not None else None) == expected, resource + + +@pytest.mark.asyncio +async def test_resource_resolution_is_identity_not_ip_filtered_access(): + """The resolver decides which server a resource NAMES; per-IP visibility filtering + belongs to the MCP routes and grant intersection. Filtering here would mint an + entitlement-wide unscoped bearer exactly when the caller asked to narrow, and IP drift + between authorize and token would turn a matching redemption into invalid_target.""" + from unittest.mock import patch + + from litellm.proxy._experimental.mcp_server.gateway_dcr_flow import resolve_scoped_resource_server + + with patch(_MANAGER_PATCH) as manager: + manager.get_mcp_server_by_name.return_value = _scoped_mcp_server() + result = resolve_scoped_resource_server(_request(), SCOPED_RESOURCE) + assert result is not None + manager.get_mcp_server_by_name.assert_called_once_with("github") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py index 3e934577a66..f25d3baea0a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -137,7 +137,7 @@ def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): [ ("/mcp/sample_docs", ["sample_docs"]), # Server names may contain at most one slash (mirrors - # ``_extract_target_server_names_from_path``), so when more than two + # ``extract_target_server_names_from_path``), so when more than two # segments follow ``/mcp/`` the first two are treated as the name. ("/mcp/sample_docs/tools/list", ["sample_docs/tools"]), ("/mcp/custom_solutions/user_123", ["custom_solutions/user_123"]), diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 55cc8565a71..99181f0f087 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -9859,3 +9859,66 @@ class TestToolAuthorizationIsNotConditionalOnLogging: ) upstream.assert_awaited_once() + + +class TestSessionResourceScopeIntersect: + """LIT-4917: the sealed session scope intersects the admitted subject's resolved server + set at the single convergence point every fan-out and tool call reads, covering the + exception fallback so a resolver fault never widens a scoped bearer.""" + + def _admitted_auth(self, scope): + from litellm.proxy._types import UserAPIKeyAuth + + auth = UserAPIKeyAuth(user_id="scoped-user") + auth.mcp_admitted_user_subject = True + auth.mcp_session_resource_server_id = scope + return auth + + def test_scope_reader_is_none_for_keys_and_unscoped_subjects(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import UserAPIKeyAuth + + assert MCPServerManager._admitted_session_resource_scope(None) is None + assert MCPServerManager._admitted_session_resource_scope(UserAPIKeyAuth(user_id="u")) is None + assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth(None)) is None + + def test_scope_reader_returns_sealed_scope_for_admitted_subjects(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth("b")) == "b" + + @pytest.mark.asyncio + async def test_get_allowed_mcp_servers_scopes_past_operator_open_union(self): + """The intersect applies AFTER the operator-open (allow_all_keys) union, so a scoped + bearer cannot reach an allow-all server outside its scope, and applies on the + exception fallback so a resolver fault yields the scoped subset of allow-all rather + than the whole set.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + manager = MCPServerManager() + auth = self._admitted_auth("granted-id") + with ( + patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=["granted-id", "other-id"], + ), + patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]), + ): + allowed = await manager.get_allowed_mcp_servers(auth) + assert allowed == ["granted-id"] + + with ( + patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=["open-id", "granted-id"]), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler.get_allowed_mcp_servers", + new_callable=AsyncMock, + side_effect=RuntimeError("resolver down"), + ), + patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]), + ): + fallback = await manager.get_allowed_mcp_servers(auth) + assert fallback == ["granted-id"] From d0be6eee8a34cb2a9c62a63a5aca71bdbd18f5db Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 15:22:06 -0700 Subject: [PATCH 16/19] fix(passthrough): stop forwarding client Accept-Encoding upstream --- litellm/passthrough/utils.py | 3 +++ .../test_vertex_passthrough_load_balancing.py | 25 +++++++++++++++++++ 2 files changed, 28 insertions(+) diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index e419322dca6..1572913f46e 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -69,6 +69,9 @@ class BasePassthroughUtils: # Header We Should NOT forward request_headers.pop("content-length", None) request_headers.pop("host", None) + # accept-encoding must stay client-negotiated: forwarding e.g. "br" when + # the brotli package is absent relays undecodable bytes to the caller + request_headers.pop("accept-encoding", None) custom_header_names: Final = {header_name.lower() for header_name in headers} for header_name in list(request_headers.keys()): diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index aaf1dad4910..02464ad5caa 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -568,6 +568,31 @@ def test_forward_headers_custom_wins_case_insensitive_over_request_authorization assert result["x-request-id"] == "req-123" +def test_forward_headers_never_forwards_client_accept_encoding(): + """ + The client's Accept-Encoding must not reach the upstream provider: the proxy's + HTTP client decodes the upstream body and advertises only encodings it can + decode. Forwarding e.g. "br" on an install without the brotli package makes + the proxy relay raw compressed bytes with the content-encoding header stripped + (garbled JSON for /v1/models and count_tokens through the Anthropic passthrough). + """ + from litellm.passthrough.utils import BasePassthroughUtils + + request_headers = { + "accept-encoding": "gzip, deflate, br, zstd", + "x-request-id": "req-123", + } + + result = BasePassthroughUtils.forward_headers_from_request( + request_headers=request_headers, + headers={}, + forward_headers=True, + ) + + assert "accept-encoding" not in result + assert result["x-request-id"] == "req-123" + + @pytest.mark.asyncio async def test_vertex_passthrough_custom_model_name_replaced_in_url(): """ From 540caa6574eb4040f2c114236aea26de1f3afe88 Mon Sep 17 00:00:00 2001 From: tin-berri Date: Sat, 15 Aug 2026 15:30:28 -0700 Subject: [PATCH 17/19] feat(ui): direction picker and reverse-mode display for shadow evals (#36994) * feat(ui): direction picker and reverse-mode display for shadow evals * fix(ui): include configured model groups in the shadow eval baseline picker --- .../_components/ShadowEvalSection.test.tsx | 55 +++++ .../_components/ShadowEvalSection.tsx | 194 ++++++++++++++---- .../hooks/models/useModels.test.ts | 27 +++ .../app/(dashboard)/hooks/models/useModels.ts | 21 ++ 4 files changed, 253 insertions(+), 44 deletions(-) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx index d4d26650086..e342ee33f25 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx @@ -46,6 +46,7 @@ vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ { model_name: "gpt-auto", litellm_params: { model: "auto_router/gpt-auto" } }, ], })), + usePlainModelGroups: vi.fn(() => new Set(["prod-claude"])), })); vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ @@ -72,6 +73,8 @@ const job = (overrides: Partial = {}): ShadowEvalJob => ({ job_id: "job-1", status: "running", router_name: "claude-auto", + direction: "forward", + baseline_model: null, judge_model: "anthropic/claude-sonnet-5", shadow_percentage: 10, max_turns: 200, @@ -367,6 +370,58 @@ describe("ShadowEvalSection", () => { expect(start.mutate).toHaveBeenCalledWith(expectedBody); }); + it("requires a baseline model in reverse mode and submits it, while forward mode never shows the picker", async () => { + const user = userEvent.setup(); + const { start } = mockHooks({}); + render(); + + expect(screen.queryByPlaceholderText("Select a baseline model")).not.toBeInTheDocument(); + + await user.click(screen.getByText("Adoption check: key's traffic vs the router")); + await user.click(await screen.findByText("Regression check: router's picks vs a baseline")); + await user.click(screen.getByPlaceholderText("Search keys by alias")); + await user.click(await screen.findByText("prod-alpha")); + await user.click(screen.getByPlaceholderText("Select an auto-router")); + await user.click(await screen.findByText("gpt-auto")); + await user.click(screen.getByPlaceholderText("Select a judge model")); + await user.click(await screen.findByRole("option", { name: /anthropic\/claude-sonnet-5/ })); + + expect(screen.getByText("Start shadow eval")).toBeDisabled(); + + await user.click(screen.getByPlaceholderText("Select a baseline model")); + expect(await screen.findByRole("option", { name: /openai\/gpt-4o/ })).toBeInTheDocument(); + await user.click(screen.getByRole("option", { name: /prod-claude/ })); + await user.click(screen.getByText("Start shadow eval")); + + const expectedBody = { + api_key_id: "hash-alpha", + router_name: "gpt-auto", + direction: "reverse", + baseline_model: "prod-claude", + shadow_percentage: 10, + duration_days: 7, + max_turns: 200, + judge_model: "anthropic/claude-sonnet-5", + }; + expect(start.mutate).toHaveBeenCalledWith(expectedBody); + }); + + it("flips the arm labels and headline for a reverse job's results", () => { + const j = job({ direction: "reverse", baseline_model: "openai/gpt-4o" }); + mockHooks({ jobs: [j], detailsById: { "job-1": j } }); + render(); + + expect(screen.getByText(/on 10% of its traffic/)).toBeInTheDocument(); + expect(screen.getByText("Router matched or beat the baseline")).toBeInTheDocument(); + expect(screen.getByText("52.0%")).toBeInTheDocument(); + expect(screen.getByText(/Router won 30.0%/)).toBeInTheDocument(); + expect(screen.getByText(/Baseline won 48.0%/)).toBeInTheDocument(); + expect(screen.getAllByText("Baseline wins")).toHaveLength(2); + expect(screen.getByText("Router pick")).toBeInTheDocument(); + expect(screen.queryByText(/Current model/)).not.toBeInTheDocument(); + expect(screen.queryByText("Compared against")).not.toBeInTheDocument(); + }); + it("keeps an older job's verdicts reachable through the previous evaluations list", async () => { const user = userEvent.setup(); const emptyOverrides: Partial = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx index 711fc1af539..df2989d3990 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx @@ -5,7 +5,7 @@ import React, { useMemo, useState } from "react"; import { useInfiniteKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; -import { useAutoRouters } from "@/app/(dashboard)/hooks/models/useModels"; +import { useAutoRouters, usePlainModelGroups } from "@/app/(dashboard)/hooks/models/useModels"; import { PaginatedSearchSelect } from "@/components/shared/PaginatedSearchSelect"; import { SearchSelect, type SearchSelectOption } from "@/components/shared/SearchSelect"; import { Badge } from "@/components/ui/badge"; @@ -31,6 +31,37 @@ const pct = (value: number): string => `${value.toFixed(1)}%`; const MIN_TURNS_FOR_CONFIDENCE = 30; +type ShadowEvalDirection = ShadowEvalJob["direction"]; + +const otherArmLabel = (direction: ShadowEvalDirection): string => + direction === "reverse" ? "Baseline" : "Current model"; + +const routerWinRate = (direction: ShadowEvalDirection, slice: ShadowEvalSlice): number => + direction === "reverse" ? slice.real_win_rate_pct : slice.shadow_win_rate_pct; + +const otherArmWinRate = (direction: ShadowEvalDirection, slice: ShadowEvalSlice): number => + direction === "reverse" ? slice.shadow_win_rate_pct : slice.real_win_rate_pct; + +const routerMatchedOrBeatPct = ( + direction: ShadowEvalDirection, + results: NonNullable, +): number => + direction === "reverse" + ? 100 - results.overall_shadow_win_rate_pct + : results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct; + +const jobHeadline = (job: ShadowEvalJob): React.ReactNode => + job.direction === "reverse" ? ( + <> + Comparing {job.router_name} to{" "} + {job.baseline_model} on {job.shadow_percentage}% of its traffic + + ) : ( + <> + Shadowing {job.shadow_percentage}% via {job.router_name} + + ); + const isActive = (job: ShadowEvalJob): boolean => job.status === "running"; const endsIn = (endsAt: string | null | undefined): string | null => { @@ -54,16 +85,22 @@ const StatusBadge: React.FC<{ status: string }> = ({ status }) => ( ); -const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSlice[] }> = ({ groupHeader, slices }) => ( +const SliceTable: React.FC<{ + groupHeader: string; + direction: ShadowEvalDirection; + slices: readonly ShadowEvalSlice[]; +}> = ({ groupHeader, direction, slices }) => ( {groupHeader} - {["Judged turns", "Router wins", "Current model wins", "Ties", "Judge confidence"].map((label) => ( - - {label} - - ))} + {["Judged turns", "Router wins", `${otherArmLabel(direction)} wins`, "Ties", "Judge confidence"].map( + (label) => ( + + {label} + + ), + )} @@ -77,9 +114,9 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli {slice.turn_count.toLocaleString()} - {pct(slice.shadow_win_rate_pct)} + {pct(routerWinRate(direction, slice))} - {pct(slice.real_win_rate_pct)} + {pct(otherArmWinRate(direction, slice))} {pct(slice.tie_rate_pct)} {slice.avg_judge_confidence.toFixed(2)} @@ -88,13 +125,23 @@ const SliceTable: React.FC<{ groupHeader: string; slices: readonly ShadowEvalSli
); -const VerdictBar: React.FC<{ results: NonNullable }> = ({ results }) => { - const routerWins = results.overall_shadow_win_rate_pct; +const VerdictBar: React.FC<{ direction: ShadowEvalDirection; results: NonNullable }> = ({ + direction, + results, +}) => { const ties = results.overall_tie_rate_pct; + const routerWins = + direction === "reverse" + ? Math.max(0, 100 - results.overall_shadow_win_rate_pct - ties) + : results.overall_shadow_win_rate_pct; const segments = [ { label: "Router won", value: routerWins, fill: "bg-emerald-500" }, { label: "Tie", value: ties, fill: "bg-emerald-200" }, - { label: "Current model won", value: Math.max(0, 100 - routerWins - ties), fill: "bg-muted-foreground/30" }, + { + label: `${otherArmLabel(direction)} won`, + value: Math.max(0, 100 - routerWins - ties), + fill: "bg-muted-foreground/30", + }, ]; return (
@@ -133,20 +180,22 @@ const ResultsBody: React.FC<{ job: ShadowEvalJob; resultsError?: boolean }> = ({ <>

- Router matched or beat your current model -

-

- {pct(results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct)} + Router matched or beat {job.direction === "reverse" ? "the baseline" : "your current model"}

+

{pct(routerMatchedOrBeatPct(job.direction, results))}

of {(job.judged_count ?? 0).toLocaleString()} judged responses

- + {results.by_current_model.length > 0 && ( - + )} {results.by_tier.length > 0 && (
0 ? "border-t" : ""}> - +
)} @@ -168,9 +217,7 @@ const JobResults: React.FC<{
-

- Shadowing {job.shadow_percentage}% via {job.router_name} -

+

{jobHeadline(job)}

{(job.judged_count ?? 0).toLocaleString()} of {job.max_turns.toLocaleString()} turns judged ·{" "} {(job.error_count ?? 0).toLocaleString()} errored · {usd(job.judge_spend ?? 0)} judge spend @@ -201,25 +248,55 @@ interface CostMapEntry { mode?: string; } -const useJudgeModelOptions = (): SearchSelectOption[] => { +const useChatModelNames = (): string[] => { const { data: costMap } = useModelCostMap(); + return useMemo(() => { + if (!costMap) return []; + const chatModels = Object.entries(costMap as Record) + .filter(([, value]) => value?.mode === "chat" && value?.litellm_provider) + .map(([key, value]) => (key.startsWith(`${value.litellm_provider}/`) ? key : `${value.litellm_provider}/${key}`)); + return [...new Set(chatModels)].toSorted((a, b) => a.localeCompare(b)); + }, [costMap]); +}; + +const useJudgeModelOptions = (): SearchSelectOption[] => { + const chatModels = useChatModelNames(); return useMemo(() => { const pinned: SearchSelectOption[] = RECOMMENDED_JUDGE_MODELS.map((model) => ({ label: model, value: model, sublabel: "Recommended", })); - if (!costMap) return pinned; const pinnedNames = new Set(RECOMMENDED_JUDGE_MODELS); - const chatModels = Object.entries(costMap as Record) - .filter(([, value]) => value?.mode === "chat" && value?.litellm_provider) - .map(([key, value]) => (key.startsWith(`${value.litellm_provider}/`) ? key : `${value.litellm_provider}/${key}`)); - const rest = [...new Set(chatModels)] - .filter((model) => !pinnedNames.has(model)) - .toSorted((a, b) => a.localeCompare(b)) - .map((model) => ({ label: model, value: model })); + const rest = chatModels.filter((model) => !pinnedNames.has(model)).map((model) => ({ label: model, value: model })); return [...pinned, ...rest]; - }, [costMap]); + }, [chatModels]); +}; + +const useBaselineModelOptions = (): SearchSelectOption[] => { + const configuredGroups = usePlainModelGroups(); + const chatModels = useChatModelNames(); + return useMemo(() => { + const configured = [...configuredGroups] + .toSorted((a, b) => a.localeCompare(b)) + .map((model) => ({ label: model, value: model, sublabel: "Configured on this gateway" })); + const rest = chatModels + .filter((model) => !configuredGroups.has(model)) + .map((model) => ({ label: model, value: model })); + return [...configured, ...rest]; + }, [configuredGroups, chatModels]); +}; + +const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[] = [ + { value: "forward", label: "Adoption check: key's traffic vs the router" }, + { value: "reverse", label: "Regression check: router's picks vs a baseline" }, +] as const; + +const START_FORM_DESCRIPTION: Record = { + forward: + "Duplicates a sampled slice of the key's traffic through the auto-router and has an LLM judge compare both answers blind. The router's answers are never served to users; judge calls bill to the shadowed key.", + reverse: + "Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. The baseline's answers are never served to users; judge calls bill to the shadowed key.", }; const DURATION_OPTIONS = [ @@ -282,12 +359,15 @@ const StartForm: React.FC = () => { const { accessToken } = useAuthorized(); const [apiKeyId, setApiKeyId] = useState(""); const [routerName, setRouterName] = useState(""); + const [direction, setDirection] = useState("forward"); + const [baselineModel, setBaselineModel] = useState(""); const [percentage, setPercentage] = useState("10"); const [durationDays, setDurationDays] = useState("7"); const [judgeModel, setJudgeModel] = useState(""); const [maxTurns, setMaxTurns] = useState("200"); const { data: autoRouters } = useAutoRouters(); const judgeModelOptions = useJudgeModelOptions(); + const baselineModelOptions = useBaselineModelOptions(); const start = useStartShadowEval(); const routerOptions = useMemo(() => { @@ -301,14 +381,17 @@ const StartForm: React.FC = () => { const percentageValid = parsedPct >= 0.1 && parsedPct <= 100; const parsedMaxTurns = Number.parseInt(maxTurns, 10); const maxTurnsValid = parsedMaxTurns >= 1 && parsedMaxTurns <= 2000; - const filled = [apiKeyId, routerName, judgeModel].every((field) => field !== ""); + const filled = + [apiKeyId, routerName, judgeModel].every((field) => field !== "") && + (direction === "forward" || baselineModel !== ""); const boundsValid = percentageValid && maxTurnsValid; const valid = Boolean(accessToken) && filled && boundsValid; const handleStart = () => { const startBody = { api_key_id: apiKeyId, router_name: routerName, - direction: "forward" as const, + direction, + ...(direction === "reverse" ? { baseline_model: baselineModel } : {}), shadow_percentage: parsedPct, duration_days: Number.parseInt(durationDays, 10), max_turns: parsedMaxTurns, @@ -321,13 +404,27 @@ const StartForm: React.FC = () => { Start a shadow eval -

- Duplicates a sampled slice of the key's traffic through the auto-router and has an LLM judge compare both - answers blind. The router's answers are never served to users; judge calls bill to the shadowed key. -

+

{START_FORM_DESCRIPTION[direction]}

+ + + @@ -390,6 +487,17 @@ const StartForm: React.FC = () => {

Enter a value from 1 to 2000

)} + {direction === "reverse" && ( + + + + )} { const previousSummary = (job: ShadowEvalJob): string => { const results = job.results; - if (results) return pct(results.overall_shadow_win_rate_pct + results.overall_tie_rate_pct); + if (results) return pct(routerMatchedOrBeatPct(job.direction, results)); return job.judged_count === 0 ? "no verdicts" : "view results"; }; @@ -429,9 +537,7 @@ const PreviousJob: React.FC<{ job: ShadowEvalJob }> = ({ job }) => {
-

- {shown.shadow_percentage}% via {shown.router_name} -

+

{jobHeadline(shown)}

{shown.judged_count != null && `${shown.judged_count.toLocaleString()} judged · ${(shown.error_count ?? 0).toLocaleString()} errored · ${usd(shown.judge_spend ?? 0)} judge spend · `} @@ -507,8 +613,8 @@ const ShadowEvalSection: React.FC = () => {

Shadow eval

- Would the auto-router have answered as well as the models you use today? Find out on your real traffic, before - switching anything. + Blind-judge the auto-router on your real traffic: against the models a key uses today before switching, or + against a fixed baseline after it has switched.

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts index f04c4b7bfcd..411e8402e11 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.test.ts @@ -5,6 +5,7 @@ import { beforeEach, describe, expect, it, vi } from "vitest"; import { isAutoRouterDeployment, selectAutoRouterModelGroups, + selectPlainModelGroups, useAllProxyModels, useAutoRouterModelGroups, useAutoRouters, @@ -977,6 +978,32 @@ describe("selectAutoRouterModelGroups", () => { }); }); +describe("selectPlainModelGroups", () => { + it("keeps only non-auto-router model groups", () => { + const deployments: AutoRouterCandidateDeployment[] = [ + { model_name: "smart-router", litellm_params: { model: "auto_router/complexity_router" } }, + { model_name: "claude-haiku", litellm_params: { model: "anthropic/claude-haiku-4-5" } }, + { model_name: "claude-sonnet", litellm_params: { model: "anthropic/claude-sonnet-4-5" } }, + { model_name: "cheap-router", litellm_params: { model: "auto_router/adaptive_router" } }, + ]; + + expect(selectPlainModelGroups(deployments)).toEqual(new Set(["claude-haiku", "claude-sonnet"])); + }); + + it("drops a group name that also fronts an auto-router deployment", () => { + const deployments: AutoRouterCandidateDeployment[] = [ + { model_name: "shared-name", litellm_params: { model: "auto_router/complexity_router" } }, + { model_name: "shared-name", litellm_params: { model: "anthropic/claude-sonnet-4-5" } }, + ]; + + expect(selectPlainModelGroups(deployments)).toEqual(new Set()); + }); + + it("drops deployments that have no public model_name", () => { + expect(selectPlainModelGroups([{ model_name: "", litellm_params: { model: "openai/gpt-4o" } }])).toEqual(new Set()); + }); +}); + describe("useAutoRouterModelGroups", () => { let queryClient: QueryClient; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts index 68044a45edc..a5fbc433ea3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/models/useModels.ts @@ -123,6 +123,16 @@ export const selectAutoRouterModelGroups = (deployments: AutoRouterCandidateDepl export const selectAutoRouterDeployments = (deployments: AutoRouterDeployment[]): AutoRouterDeployment[] => deployments.filter(isAutoRouterDeployment); +export const selectPlainModelGroups = (deployments: AutoRouterCandidateDeployment[]): ReadonlySet => { + const autoRouterGroups = selectAutoRouterModelGroups(deployments); + return new Set( + deployments + .map((deployment) => deployment.model_name) + .filter((modelName): modelName is string => Boolean(modelName)) + .filter((modelName) => !autoRouterGroups.has(modelName)), + ); +}; + export const fetchAllModelDeployments = async ( accessToken: string, userId: string, @@ -172,6 +182,17 @@ export const useAutoRouterModelGroups = (): ReadonlySet => { return data ?? NO_AUTO_ROUTERS; }; +export const usePlainModelGroups = (): ReadonlySet => { + const { accessToken, userId, userRole } = useAuthorized(); + const { data } = useQuery>({ + queryKey: autoRouterListKey(userId, userRole), + queryFn: async () => await fetchAllModelDeployments(accessToken!, userId!, userRole!), + enabled: Boolean(accessToken && userId && userRole), + select: selectPlainModelGroups, + }); + return data ?? NO_AUTO_ROUTERS; +}; + export const useAutoRouters = (): UseQueryResult => { const { accessToken, userId, userRole } = useAuthorized(); return useQuery({ From d7d10be0639395da813d6b9c88933bfafbcb3f85 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 15 Aug 2026 15:31:23 -0700 Subject: [PATCH 18/19] fix(guardrails): return the full PANW AIRS scan response on blocked requests (#37036) * fix(guardrails): return the full PANW AIRS scan response on blocked requests The blocked-request error detail was assembled from a hardcoded allowlist, so audit fields like prompt_detection_details, prompt_masked_data, source, transaction_id and session_id never reached the client even though AIRS returned them. Resolves LIT-5638 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(guardrails): drop redundant comment in AIRS error detail Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(panw_prisma_airs): withhold response_masked_data from the blocked-response error The full AIRS passthrough also reached the response-side block path, where response_masked_data carries the model's own generation. That branch is only reached when mask_response_content is False, so the operator had explicitly declined to deliver that text, and the error body handed it back anyway. Withhold response_masked_data from the client-visible detail. prompt_masked_data stays: it is the caller's own input and one of the fields the ticket asks for. Every other AIRS field, including prompt_detection_details, source, transaction_id and session_id, is unchanged. * fix(panw_prisma_airs): withhold generated tool args from response-side blocks _scan_tool_calls_for_guardrail calls AIRS with is_response=False because tool_event is request-side in the AIRS schema, so AIRS returns the scanned tool arguments under prompt_masked_data. When the tool calls being scanned are the model's own output, that key holds generated content, and the _CLIENT_HIDDEN_SCAN_FIELDS default (response_masked_data, empty on this path) does not cover it. With the default mask_response_content=False the block branch then shipped the model's masked tool arguments in the 400 -- the same content channel this PR closed for response_masked_data. _build_error_detail takes an extra_hidden_fields argument so the withholding stays in one place, and the tool-call block branch passes prompt_masked_data when is_response is True. Request-side blocks are unchanged and still carry prompt_masked_data, which is what LIT-5638 asks for. Co-Authored-By: Claude Opus 5 (1M context) * style(panw_prisma_airs): apply ruff format Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: Yucheng Zhu Co-authored-by: Claude Opus 5 (1M context) --- .../panw_prisma_airs/panw_prisma_airs.py | 60 +++-- .../guardrail_hooks/test_panw_prisma_airs.py | 229 ++++++++++++++++++ 2 files changed, 266 insertions(+), 23 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index 3765771247d..7e641814deb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -71,6 +71,14 @@ class PanwPrismaAirsHandler(CustomGuardrail): _PROVIDER_NAME = "panw_prisma_airs" + #: AIRS fields withheld from the client-visible error detail. + #: ``response_masked_data`` is the model's own generation. The block branch that builds + #: this detail is only reached when ``mask_response_content`` is False, so echoing it + #: back would hand the caller exactly the text the operator declined to deliver. + #: ``prompt_masked_data`` is deliberately NOT withheld: it is the caller's own input, + #: and it is one of the fields the ticket asks for. + _CLIENT_HIDDEN_SCAN_FIELDS: Final = frozenset({"response_masked_data"}) + def __init__( self, guardrail_name: str, @@ -632,12 +640,21 @@ class PanwPrismaAirsHandler(CustomGuardrail): choice.message.function_call.arguments = masked_text def _build_error_detail( - self, scan_result: Mapping[str, object], is_response: bool = False + self, + scan_result: Mapping[str, object], + is_response: bool = False, + also_hide: str | None = None, ) -> Mapping[str, Mapping[str, object]]: - """Build enhanced error detail with scan information.""" + """Build enhanced error detail with scan information. + + ``also_hide`` names one more scan field to withhold, for the caller that knows + its AIRS verdict carries model-generated content under a key that is normally + caller input. + """ action_type: Final = "Response" if is_response else "Prompt" code_suffix: Final = "_response_blocked" if is_response else "_blocked" - detection_key: Final = "response_detected" if is_response else "prompt_detected" + + hidden_fields: Final = self._CLIENT_HIDDEN_SCAN_FIELDS.union(() if also_hide is None else (also_hide,)) category: Final = scan_result.get("category", "unknown") default_msg: Final = f"{action_type} blocked by PANW Prisma AI Security policy (Category: {category})" @@ -653,8 +670,13 @@ class PanwPrismaAirsHandler(CustomGuardrail): }, ) - error_detail: Final[dict[str, dict[str, object]]] = { + return { "error": { + **{ + key: value + for key, value in scan_result.items() + if not key.startswith("_") and key not in hidden_fields + }, "message": error_msg, "type": "guardrail_violation", "code": f"panw_prisma_airs{code_suffix}", @@ -663,24 +685,6 @@ class PanwPrismaAirsHandler(CustomGuardrail): } } - # Add optional fields if present - optional_fields: Final = [ - "scan_id", - "report_id", - "profile_name", - "profile_id", - "tr_id", - ] - for field in optional_fields: - if scan_result.get(field): - error_detail["error"][field] = scan_result[field] - - # Add detection details - if scan_result.get(detection_key): - error_detail["error"][detection_key] = scan_result[detection_key] - - return error_detail - def _record_scan_id(self, request_data: dict[str, Any], scan_result: Mapping[str, object]) -> None: """Surface the AIRS scan id on the response, so allowed calls are auditable too.""" scan_id: Final = scan_result.get("scan_id") @@ -1481,7 +1485,17 @@ class PanwPrismaAirsHandler(CustomGuardrail): ): self._set_tool_call_arguments(tool_call, masked_text) else: - error_detail = self._build_error_detail(scan_result, is_response=is_response) + # tool_event scans are request-side in the AIRS schema, so AIRS returns + # the model's own tool arguments under prompt_masked_data. On a + # response-side block that is generated content, not caller input, and + # the class-level default only withholds response_masked_data — which is + # empty on this path. Withhold it explicitly so the 400 does not become + # the content channel this branch declined to deliver. + error_detail = self._build_error_detail( + scan_result, + is_response=is_response, + also_hide="prompt_masked_data" if is_response else None, + ) raise HTTPException(status_code=400, detail=error_detail) @staticmethod diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py index 17d4a3e304a..fc1465c7c14 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_panw_prisma_airs.py @@ -5647,6 +5647,235 @@ class TestPanwAirsScanIdExposure: assert "guardrail_scan_ids" in _UNTRUSTED_METADATA_CONTROL_FIELDS assert "guardrail_scan_ids" in _UNTRUSTED_ROOT_CONTROL_FIELDS +class TestPanwAirsBlockedErrorDetailPassthrough: + """Regression tests for the full AIRS scan response on blocks. + + Before the fix, the error detail was built from a hardcoded allowlist + (scan_id, report_id, profile_name, profile_id, tr_id, prompt/response_detected), + so audit-relevant fields such as prompt_detection_details, prompt_masked_data, + source, transaction_id and session_id never reached the client. + """ + + _FULL_BLOCK_RESPONSE = { + "action": "block", + "category": "malicious", + "scan_id": "b2f0a4be-1f6f-4f9a-9f3d-4b6a9d8b1c0e", + "report_id": "R0000000000000000000", + "tr_id": "test-call-id", + "profile_id": "6f5c9f6e-2d0b-4d3f-8a1e-9b7c5d4e3f2a", + "profile_name": "test_profile", + "source": "prisma_airs", + "transaction_id": "4b8c1e2f-5a6d-4c3b-9e8f-1a2b3c4d5e6f", + "session_id": "3a2b1c0d-9e8f-4a7b-8c6d-5e4f3a2b1c0d", + "timeout": False, + "errors": [], + "prompt_detected": {"dlp": True, "injection": False, "url_cats": False}, + "prompt_detection_details": { + "dlp_report": { + "dlp_report_id": "1234567890", + "dlp_profile_name": "Sensitive Content", + "data_pattern_rule1_verdict": "MATCHED", + } + }, + "prompt_masked_data": {"data": "my ssn is XXX-XX-XXXX"}, + "response_detected": {"dlp": False, "url_cats": False}, + "response_detection_details": {}, + "response_masked_data": {}, + } + + @pytest.mark.asyncio + @pytest.mark.parametrize("is_response", [False, True]) + async def test_block_returns_every_airs_field( + self, base_handler, user_api_key_dict, safe_prompt_data, is_response + ): + response = ModelResponse( + id="test_id", + choices=[ + Choices(index=0, message=Message(role="assistant", content="Test response")), + ], + model="gpt-3.5-turbo", + ) + + with patch.object( + base_handler, "_call_panw_api", return_value=copy.deepcopy(self._FULL_BLOCK_RESPONSE) + ): + with pytest.raises(HTTPException) as exc_info: + if is_response: + await base_handler.async_post_call_success_hook( + data=safe_prompt_data, + user_api_key_dict=user_api_key_dict, + response=response, + ) + else: + await base_handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=None, + data=safe_prompt_data, + call_type="completion", + ) + + error = exc_info.value.detail["error"] + for field, value in self._FULL_BLOCK_RESPONSE.items(): + if field == "category": + continue + if field in PanwPrismaAirsHandler._CLIENT_HIDDEN_SCAN_FIELDS: + # Withheld on purpose, covered by TestPanwAirsErrorDetailWithheldFields + continue + assert error[field] == value, f"{field} missing or altered in blocked-request error" + + assert error["category"] == "malicious" + assert error["type"] == "guardrail_violation" + assert error["guardrail"] == "test_panw_airs" + assert error["code"] == ("panw_prisma_airs_response_blocked" if is_response else "panw_prisma_airs_blocked") + assert "PANW Prisma AI Security policy" in error["message"] + + def test_internal_control_flags_are_not_leaked(self, base_handler): + detail = base_handler._build_error_detail( + { + "action": "block", + "category": "malicious", + "scan_id": "scan-1", + "_always_block": True, + "_is_transient": True, + } + ) + + assert "_always_block" not in detail["error"] + assert "_is_transient" not in detail["error"] + assert detail["error"]["scan_id"] == "scan-1" + + +class TestPanwAirsErrorDetailWithheldFields: + """The blocked-request passthrough must not become a content channel. + + ``response_masked_data`` is the model's own generation. The block branch is only + reached when ``mask_response_content`` is False, so echoing it back would hand the + caller exactly the text the operator declined to deliver. ``error`` is AIRS's own + message about the operator's Strata Cloud Manager profile configuration. + + ``prompt_masked_data`` is deliberately NOT withheld by default: it is the caller's + own input, and it is one of the fields LIT-5638 asks for. The one exception is the + response-side tool-call path, covered by + ``TestPanwAirsToolCallBlockWithholdsGeneratedArgs`` below — tool_event scans are + request-side in the AIRS schema, so there the key holds model output instead. + """ + + @pytest.mark.parametrize("is_response", [False, True]) + def test_response_masked_data_never_reaches_client(self, base_handler, is_response): + detail = base_handler._build_error_detail( + { + "action": "block", + "category": "sensitive_data", + "scan_id": "scan-1", + "response_detected": {"dlp": True}, + "response_masked_data": {"data": "routing number XXXXXXXXXX"}, + "prompt_masked_data": {"data": "my ssn is XXX-XX-XXXX"}, + "prompt_detection_details": {"dlp_report": {"dlp_report_id": "1"}}, + }, + is_response=is_response, + ) + error = detail["error"] + + assert "response_masked_data" not in error + assert "routing number" not in str(error) + + # The audit fields LIT-5638 asks for still come through untouched. + assert error["scan_id"] == "scan-1" + assert error["response_detected"] == {"dlp": True} + assert error["prompt_masked_data"] == {"data": "my ssn is XXX-XX-XXXX"} + assert error["prompt_detection_details"] == {"dlp_report": {"dlp_report_id": "1"}} + + def test_upstream_airs_error_field_still_passes_through(self, base_handler): + """A 2xx AIRS body can carry its own ``error`` (see _call_panw_api's + profile-misconfiguration branch, which only logs and then blocks). It is + diagnostic rather than content, so it stays in the passthrough.""" + detail = base_handler._build_error_detail( + { + "action": "block", + "category": "malicious", + "scan_id": "scan-2", + "error": "profile not found", + } + ) + + assert detail["error"]["error"] == "profile not found" + assert detail["error"]["scan_id"] == "scan-2" + + +class TestPanwAirsToolCallBlockWithholdsGeneratedArgs: + """A response-side tool-call block must not ship the model's tool arguments. + + ``_scan_tool_calls_for_guardrail`` calls AIRS with ``is_response=False`` because + tool_event is request-side in the AIRS schema, so AIRS returns the scanned tool + arguments under ``prompt_masked_data``. When the tool calls being scanned are the + model's own output, that key holds generated content, and the class-level + ``_CLIENT_HIDDEN_SCAN_FIELDS`` default (``response_masked_data``, empty on this + path) does not cover it. + """ + + MASKED_ARGS = '{"to_account": "XXXXXXXXXX", "amount": 5000}' + + SCAN_RESULT = { + "action": "block", + "category": "sensitive_data", + "scan_id": "scan-tool-1", + "prompt_detected": {"dlp": True}, + "prompt_masked_data": {"data": MASKED_ARGS}, + "response_masked_data": {}, + } + + @staticmethod + def _tool_call(): + return ChatCompletionMessageToolCall( + id="call_1", + type="function", + function=Function( + name="transfer_funds", + arguments='{"to_account": "ACME-VENDOR-001", "amount": 5000}', + ), + ) + + async def _block(self, handler, is_response): + with patch.object(handler, "_call_panw_api", new_callable=AsyncMock) as mock_api: + mock_api.return_value = dict(self.SCAN_RESULT) + with pytest.raises(HTTPException) as exc_info: + await handler._scan_tool_calls_for_guardrail( + tool_calls=[self._tool_call()], + is_response=is_response, + metadata={}, + call_id="test-call-id", + request_data={"metadata": {}}, + start_time=datetime.now(), + ) + return exc_info.value + + @pytest.mark.asyncio + async def test_response_side_block_withholds_generated_tool_args(self): + handler = make_handler(mask_response_content=False) + # The block branch is only reached with masking off; guard the premise. + assert handler.mask_response_content is False + + exc = await self._block(handler, is_response=True) + error = exc.detail["error"] + + assert exc.status_code == 400 + assert "prompt_masked_data" not in error + assert self.MASKED_ARGS not in str(error) + + # The audit fields LIT-5638 asks for are unaffected. + assert error["scan_id"] == "scan-tool-1" + assert error["prompt_detected"] == {"dlp": True} + + @pytest.mark.asyncio + async def test_request_side_block_still_returns_masked_tool_args(self): + """Caller-supplied tool arguments stay in the verdict — that is the ticket's ask.""" + handler = make_handler(mask_request_content=False) + + exc = await self._block(handler, is_response=False) + error = exc.detail["error"] + + assert error["prompt_masked_data"] == {"data": self.MASKED_ARGS} + assert error["scan_id"] == "scan-tool-1" if __name__ == "__main__": From 90493a217f06b21cc31ee1683ebe746139cb256e Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 15 Aug 2026 15:34:28 -0700 Subject: [PATCH 19/19] fix(passthrough): protect accept-encoding from x-pass- forwarding --- litellm/passthrough/utils.py | 1 + .../test_vertex_passthrough_load_balancing.py | 1 + 2 files changed, 2 insertions(+) diff --git a/litellm/passthrough/utils.py b/litellm/passthrough/utils.py index 1572913f46e..df39b8fad48 100644 --- a/litellm/passthrough/utils.py +++ b/litellm/passthrough/utils.py @@ -18,6 +18,7 @@ _PASS_THROUGH_PROTECTED_HEADERS: Final[frozenset] = frozenset( "x-goog-api-key", "host", "content-length", + "accept-encoding", } ) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py index 02464ad5caa..8e973fc3771 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_vertex_passthrough_load_balancing.py @@ -580,6 +580,7 @@ def test_forward_headers_never_forwards_client_accept_encoding(): request_headers = { "accept-encoding": "gzip, deflate, br, zstd", + "x-pass-accept-encoding": "br", "x-request-id": "req-123", }