diff --git a/litellm/llms/anthropic/batches/transformation.py b/litellm/llms/anthropic/batches/transformation.py index e71ecd3b99c..e75cb432636 100644 --- a/litellm/llms/anthropic/batches/transformation.py +++ b/litellm/llms/anthropic/batches/transformation.py @@ -11,7 +11,7 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest from litellm.types.utils import LiteLLMBatch, LlmProviders, ModelResponse -from ..common_utils import merge_anthropic_beta_headers +from ..common_utils import merge_anthropic_beta_headers, without_caller_credential_headers if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -59,15 +59,16 @@ class AnthropicBatchesConfig(BaseBatchesConfig): merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")), "message-batches-2024-09-24", ) - _headers: Final = { + # The deployment's own credential is applied below, so a caller-supplied one must not + # ride along: without this a minted federation Bearer travels beside the caller's x-api-key. + return { + **without_caller_credential_headers(headers), "accept": "application/json", "anthropic-version": "2023-06-01", "content-type": "application/json", + **auth_header, + "anthropic-beta": merged_beta, } - _headers.update(auth_header) - headers.update(_headers) - headers["anthropic-beta"] = merged_beta - return headers def get_complete_batch_url( self, diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 3770dc66f0e..470a9addfba 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -92,9 +92,20 @@ def is_anthropic_oauth_key(value: str | None) -> bool: return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX) -def merge_anthropic_beta_headers(existing: str | None, new_beta: str | None) -> str: - """Merge comma-separated anthropic-beta header values, deduplicated and sorted.""" - betas: Final = {b.strip() for value in (existing, new_beta) if value for b in value.split(",") if b.strip()} +def merge_anthropic_beta_headers(existing: str | Sequence[str] | None, new_beta: str | Sequence[str] | None) -> str: + """Merge anthropic-beta header values, deduplicated and sorted. + + Either side may arrive as a list rather than a comma-separated string: the Skills surface + accepted a list-valued header before it shared this helper, and callers still send one. + """ + values: Final = ( + entry + for side in (existing, new_beta) + if side + for entry in ((side,) if isinstance(side, str) else side) + if isinstance(entry, str) + ) + betas: Final = {b.strip() for value in values for b in value.split(",") if b.strip()} return ",".join(sorted(betas)) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 6ee43928320..b64839423c2 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -187,9 +187,10 @@ def _allow_model_level_clientside_configurable_parameters( # ``extra_body.aws_web_identity_token``) without re-validating, so the # banned-key check has to descend into it the same way it descends into # ``litellm_embedding_config``. -_ANTHROPIC_WIF_UNCONDITIONAL_BANNED: Final[tuple[str, ...]] = tuple( - sorted(p for p in anthropic_wif_litellm_params if p != "anthropic_workspace_id") -) +_ANTHROPIC_WIF_UNCONDITIONAL_BANNED: Final[tuple[str, ...]] = anthropic_wif_litellm_params +# The Bedrock Claude Platform route reads a workspace from workspace_id or aws_workspace_id as +# well, and neither is a federation parameter, so say so rather than leaving that caller stuck. +_BEDROCK_WORKSPACE_HINT: Final = " On the Bedrock Claude Platform route, pass workspace_id or aws_workspace_id instead." def reject_server_owned_wif_params(body: Mapping[str, object]) -> None: @@ -201,6 +202,7 @@ def reject_server_owned_wif_params(body: Mapping[str, object]) -> None: raise ValueError( f"Rejected Request: {param} is a server-owned workload identity federation parameter " "and cannot be set in a request body; configure it on the deployment instead." + + (_BEDROCK_WORKSPACE_HINT if param == "anthropic_workspace_id" else "") ) diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index dd6f7e0ff8d..46bf44b77c6 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -2254,6 +2254,39 @@ class TestWifHeaderContract: assert "oauth-2025-04-20" in headers["anthropic-beta"] +class TestMergeAnthropicBetaHeaders: + """The Skills surface accepted a list-valued anthropic-beta before it shared this helper, + so the helper has to keep taking one: .split() on a list is an AttributeError.""" + + def test_list_valued_existing_header_is_merged(self): + from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers + + assert merge_anthropic_beta_headers(["skills-2025-10-02", "files-api-2025-04-14"], "oauth-2025-04-20") == ( + "files-api-2025-04-14,oauth-2025-04-20,skills-2025-10-02" + ) + + def test_list_and_comma_string_forms_agree(self): + from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers + + as_list = merge_anthropic_beta_headers(["a", "b"], "c") + as_string = merge_anthropic_beta_headers("a,b", "c") + assert as_list == as_string == "a,b,c" + + def test_skills_validate_environment_accepts_a_list_header(self, monkeypatch): + """End of the regression: the Skills surface itself must not raise on the list form.""" + from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + headers = AnthropicSkillsConfig().validate_environment( + headers={"anthropic-beta": ["files-api-2025-04-14"]}, + litellm_params=None, + ) + + assert "files-api-2025-04-14" in headers["anthropic-beta"] + assert isinstance(headers["anthropic-beta"], str) + + class TestWifServerOwnedAuthHeaderStrip: """A WIF-minted token must never ride alongside a caller-supplied credential header, but that stripping must fire only when a mint actually happened.""" @@ -2316,6 +2349,31 @@ class TestWifServerOwnedAuthHeaderStrip: assert all(caller_key not in value for value in headers.values()) assert headers["user-agent"] == "caller/1.0" + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_batches_surface_strips_caller_credentials_too(self, monkeypatch, wif_engine, header_name): + """Batches builds its own headers on the create path, so it needs the same strip: the + handler's retrieve path passes none, but this entry point takes the caller's.""" + from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicBatchesConfig().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + model="claude-sonnet-4-5", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) def test_files_surface_strips_caller_credentials_too(self, monkeypatch, wif_engine, header_name): """The files surface builds its own headers, so it needs the same strip the chat surface @@ -3211,14 +3269,14 @@ class TestWifExchangeTransportHardening: class TestWifParamsAreNotClientSettable: def test_every_minting_param_is_server_owned(self): - """Each of these selects which server-side secret is read; only the inert workspace id is - left settable, because Bedrock Claude Platform already accepts that spelling.""" + """Each of these selects which server-side secret is read, or the scope it is minted for. + The workspace id was once carved out here as inert; it is not. It is the scope of the + minted org credential, and the router merges request kwargs over deployment params, so a + caller who set it picked the scope instead of the administrator.""" from litellm.proxy.auth.auth_utils import _ANTHROPIC_WIF_UNCONDITIONAL_BANNED from litellm.types.utils import anthropic_wif_litellm_params - assert set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) == set(anthropic_wif_litellm_params) - { - "anthropic_workspace_id" - } + assert set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) == set(anthropic_wif_litellm_params) class TestWifServerOwnedParamsAreUnconditional: @@ -3275,20 +3333,50 @@ class TestWifServerOwnedParamsAreUnconditional: model="claude-sonnet-5", ) - def test_workspace_id_stays_allowed(self): - """It cannot mint anything on its own, and Bedrock Claude Platform already accepts the - spelling in a request body.""" + def test_workspace_id_is_refused_from_a_request_body(self): + """Regression, proven live against Anthropic before this was closed: a caller-supplied + workspace id reached the token endpoint, which answered "workspace_id is not a well-formed + wrkspc_ tagged ID", i.e. the caller's value had become the scope of the minted credential. + router.py merges request kwargs OVER deployment params, so it also beat the configured one.""" from litellm.proxy.auth.auth_utils import is_request_body_safe - assert ( + with pytest.raises(Exception, match="server-owned workload identity federation parameter"): is_request_body_safe( request_body={"model": "claude-sonnet-5", "anthropic_workspace_id": "wrkspc_abc"}, general_settings={}, llm_router=None, model="claude-sonnet-5", ) - is True - ) + + def test_refusal_points_bedrock_callers_at_their_own_spelling(self): + """Banning this spelling must not read as "no workspace selection anywhere": the Bedrock + Claude Platform route takes workspace_id/aws_workspace_id, neither of which is a + federation parameter, so the error names them.""" + from litellm.proxy.auth.auth_utils import is_request_body_safe + + with pytest.raises(Exception, match="workspace_id or aws_workspace_id"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "anthropic_workspace_id": "wrkspc_abc"}, + general_settings={}, + llm_router=None, + model="claude-sonnet-5", + ) + + def test_bedrock_workspace_spellings_are_untouched(self): + """The Bedrock route's own spellings stay settable, which is what keeps this ban from + removing a pre-existing capability.""" + from litellm.proxy.auth.auth_utils import is_request_body_safe + + for spelling in ("workspace_id", "aws_workspace_id"): + assert ( + is_request_body_safe( + request_body={"model": "claude-sonnet-5", spelling: "wrkspc_abc"}, + general_settings={}, + llm_router=None, + model="claude-sonnet-5", + ) + is True + ) class TestWifDisabledOnClientRedirectedBase: diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index f0dbf952e8f..70c489d49bd 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -38,9 +38,7 @@ def test_every_anthropic_wif_kwarg_key_is_request_banned(): from litellm.litellm_core_utils.get_litellm_params import ANTHROPIC_WIF_KWARGS_KEYS from litellm.proxy.auth.auth_utils import _ANTHROPIC_WIF_UNCONDITIONAL_BANNED - unconditionally_bannable = ANTHROPIC_WIF_KWARGS_KEYS - {"anthropic_workspace_id"} - - assert unconditionally_bannable <= set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) + assert ANTHROPIC_WIF_KWARGS_KEYS == set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) class TestCustomAuthCommonChecksWarning: @@ -132,9 +130,7 @@ class TestGetKeyModelRpmLimit: """Should fall back to team metadata when key metadata exists but has no model_rpm_limit.""" user_api_key_dict = UserAPIKeyAuth( api_key="sk-123", - metadata={ - "some_other_key": "value" - }, # Has metadata, but not model_rpm_limit + metadata={"some_other_key": "value"}, # Has metadata, but not model_rpm_limit team_metadata={"model_rpm_limit": {"gpt-4": 50}}, ) result = get_key_model_rpm_limit(user_api_key_dict) @@ -216,9 +212,7 @@ class TestGetKeyModelTpmLimit: """Should fall back to team metadata when key metadata exists but has no model_tpm_limit.""" user_api_key_dict = UserAPIKeyAuth( api_key="sk-123", - metadata={ - "some_other_key": "value" - }, # Has metadata, but not model_tpm_limit + metadata={"some_other_key": "value"}, # Has metadata, but not model_tpm_limit team_metadata={"model_tpm_limit": {"gpt-4": 5000}}, ) result = get_key_model_tpm_limit(user_api_key_dict) @@ -329,9 +323,7 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders: request_body = {"user": "body-user"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "header-customer" def test_should_fall_back_to_body_when_no_standard_header(self): @@ -340,9 +332,7 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders: request_body = {"user": "body-user"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "body-user" @@ -384,8 +374,7 @@ def test_get_model_from_request_enforces_when_builtin_handler_dispatched(): enforced. Same request path as above, but dispatched to a non-pass-through endpoint: the model must NOT be suppressed.""" - def builtin_chat_completions(): - ... + def builtin_chat_completions(): ... assert ( get_model_from_request( @@ -504,9 +493,7 @@ def test_get_model_from_request_extracts_unified_file_id_models(): "litellm_proxy:application/octet-stream;unified_id,test-id;" "target_model_names,model-a,model-b;llm_output_file_id,file-provider-id" ) - encoded_unified_file_id = ( - base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") - ) + encoded_unified_file_id = base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") assert get_model_from_request( request_data={"file_id": encoded_unified_file_id}, @@ -566,9 +553,7 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): model_id="veo-3.1-generate-001", ) llm_router = MagicMock() - llm_router.resolve_model_name_from_model_id.return_value = ( - "gcp/google/veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001" assert ( get_model_from_request( @@ -578,9 +563,7 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): ) == "gcp/google/veo-3.1-generate-001" ) - llm_router.resolve_model_name_from_model_id.assert_called_once_with( - "veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001") _BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc" @@ -611,9 +594,7 @@ def _encode_managed_id(decoded: str) -> str: return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=") -_MANAGED_BATCH_ID = _encode_managed_id( - f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123" -) +_MANAGED_BATCH_ID = _encode_managed_id(f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123") _MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id( f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;" "llm_output_file_id:provider-file-456" @@ -689,9 +670,7 @@ def test_get_model_from_request_resolves_character_id_model_with_router(): model_id="veo-3.1-generate-001", ) llm_router = MagicMock() - llm_router.resolve_model_name_from_model_id.return_value = ( - "gcp/google/veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001" assert ( get_model_from_request( @@ -701,9 +680,7 @@ def test_get_model_from_request_resolves_character_id_model_with_router(): ) == "gcp/google/veo-3.1-generate-001" ) - llm_router.resolve_model_name_from_model_id.assert_called_once_with( - "veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001") def test_get_model_from_request_only_runs_media_decoders_for_matching_fields(): @@ -850,9 +827,7 @@ def test_abbreviate_api_key_short_key_is_fully_masked(): def test_get_customer_user_header_returns_none_when_no_customer_role(): from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping - mappings = [ - {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"} - ] + mappings = [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}] result = get_customer_user_header_from_mapping(mappings) assert result is None @@ -905,9 +880,7 @@ def test_get_end_user_id_returns_id_from_user_header_mappings(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "1234" @@ -933,9 +906,7 @@ def test_get_end_user_id_returns_first_customer_header_when_multiple_mappings_ex ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "user-456" @@ -956,9 +927,7 @@ def test_get_end_user_id_returns_none_when_no_customer_role_in_mappings(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result is None @@ -976,9 +945,7 @@ def test_get_end_user_id_falls_back_to_deprecated_user_header_name(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "user-legacy" @@ -1122,9 +1089,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1134,9 +1099,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1149,19 +1112,14 @@ class TestGetEndUserIdDropsMalformedBodyValues: """ import litellm - blob = ( - '{"device_id":"d5abe9199ee7759a","account_uuid":"",' - '"session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' - ) + blob = '{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' request_body = {"user": blob} original = litellm.validate_end_user_id_in_db litellm.validate_end_user_id_in_db = False try: with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) finally: litellm.validate_end_user_id_in_db = original @@ -1172,8 +1130,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = { "user": ( - '{"device_id":"d5abe9199ee7759a","account_uuid":"",' - '"session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' + '{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' ), } @@ -1181,9 +1138,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: litellm.validate_end_user_id_in_db = True try: with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) finally: litellm.validate_end_user_id_in_db = original @@ -1193,9 +1148,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": "alice@example.com"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1207,9 +1160,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": codex_id} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == codex_id @@ -1217,9 +1168,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": 12345} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "12345" @@ -1230,9 +1179,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1242,9 +1189,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1254,9 +1199,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1264,9 +1207,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": " ", "safety_identifier": "alice@example.com"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1285,16 +1226,12 @@ class TestGetEndUserIdDropsMalformedBodyValues: ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "alice@example.com" -def _make_deployment_dict( - model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None -) -> dict: +def _make_deployment_dict(model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None) -> dict: """Helper to build a minimal deployment dict as returned by router.get_model_list.""" litellm_params: dict = {"model": model_name} if tpm is not None: @@ -1314,9 +1251,7 @@ class TestDeploymentDefaultRpmLimit: """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 200} @@ -1328,9 +1263,7 @@ class TestDeploymentDefaultRpmLimit: metadata={"model_rpm_limit": {"model1": 10}}, ) mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 10} @@ -1350,9 +1283,7 @@ class TestDeploymentDefaultRpmLimit: """No model_name means deployment fallback is skipped.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict) assert result is None @@ -1413,9 +1344,7 @@ class TestDeploymentDefaultTpmLimit: """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 100} @@ -1427,9 +1356,7 @@ class TestDeploymentDefaultTpmLimit: metadata={"model_tpm_limit": {"model1": 20}}, ) mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 20} @@ -1449,9 +1376,7 @@ class TestDeploymentDefaultTpmLimit: """No model_name means deployment fallback is skipped.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict) assert result is None @@ -1602,7 +1527,7 @@ class TestCheckCompleteCredentialsBlocksSSRF: "litellm.proxy.auth.auth_utils.validate_url", side_effect=SSRFError(f"blocked: {blocked_url}"), ): - with pytest.raises(ValueError, match='is rejected by the SSRF guard') as exc_info: + with pytest.raises(ValueError, match="is rejected by the SSRF guard") as exc_info: check_complete_credentials( { "model": "gpt-4", @@ -1964,9 +1889,7 @@ class TestIsRequestBodySafeBlocksFallbackSmuggle: is_request_body_safe( request_body={ "model": "gpt-4", - "fallbacks": [ - {"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]} - ], + "fallbacks": [{"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]}], }, general_settings={"allow_client_side_credentials": True}, llm_router=None, @@ -1986,9 +1909,7 @@ class TestIsRequestBodySafeBlocksFallbackSmuggle: "always-fail": [ { "model": "x", - fallback_field: [ - {"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]} - ], + fallback_field: [{"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]}], } ] } @@ -2158,7 +2079,7 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields: ], ) def test_endpoint_targeting_field_in_request_body_is_rejected(self, field): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: "https://attacker.example"}, general_settings={}, @@ -2179,7 +2100,7 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields: # on the blocklist into an SSRF / credential-exfil hole. Verify # that supplying an api_key (alongside the banned param) does NOT # bypass the gate — it can only be opened by an admin opt-in. - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -2604,11 +2525,7 @@ class TestIsRequestBodySafeNestedConfig: when nested.""" with pytest.raises(ValueError, match="langfuse_host"): is_request_body_safe( - request_body={ - "litellm_embedding_config": { - "langfuse_host": "https://attacker.example.com" - } - }, + request_body={"litellm_embedding_config": {"langfuse_host": "https://attacker.example.com"}}, general_settings={}, llm_router=None, model="milvus-store", @@ -2619,11 +2536,7 @@ class TestIsRequestBodySafeNestedConfig: keep the existing escape hatch — same UX as for root-level.""" assert ( is_request_body_safe( - request_body={ - "litellm_embedding_config": { - "api_base": "https://my-azure.example.com" - } - }, + request_body={"litellm_embedding_config": {"api_base": "https://my-azure.example.com"}}, general_settings={"allow_client_side_credentials": True}, llm_router=None, model="milvus-store", @@ -2736,7 +2649,7 @@ class TestObservabilityCallbackBans: ], ) def test_observability_field_in_request_body_root_is_rejected(self, field): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: "attacker-value"}, general_settings={}, @@ -2760,13 +2673,11 @@ class TestObservabilityCallbackBans: "user_api_key_auth_metadata", ], ) - def test_observability_field_in_metadata_dict_is_rejected( - self, metadata_key, field - ): + def test_observability_field_in_metadata_dict_is_rejected(self, metadata_key, field): # Verifies the metadata walk: a value smuggled inside ``metadata`` # or ``litellm_metadata`` is just as dangerous as the same field # at the body root, and must hit the same gate. - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -2801,13 +2712,11 @@ class TestObservabilityCallbackBans: ) def test_observability_field_in_litellm_params_metadata_is_rejected(self): - with pytest.raises(ValueError, match='Rejected Request: turn_off_message_logging is not allowed') as exc: + with pytest.raises(ValueError, match="Rejected Request: turn_off_message_logging is not allowed") as exc: is_request_body_safe( request_body={ "model": "gpt-4", - "litellm_params": { - "metadata": {"turn_off_message_logging": False} - }, + "litellm_params": {"metadata": {"turn_off_message_logging": False}}, }, general_settings={}, llm_router=None, @@ -2819,22 +2728,18 @@ class TestObservabilityCallbackBans: "metadata_key", ["metadata", "litellm_metadata"], ) - def test_observability_field_in_json_string_metadata_is_rejected( - self, metadata_key - ): + def test_observability_field_in_json_string_metadata_is_rejected(self, metadata_key): # Multipart/form-data and ``extra_body`` callers send metadata as a # JSON-encoded string. The bouncer parses it before applying the # banned-params check so the JSON-string path can't smuggle past # the ``isinstance(dict)`` guard. import json - with pytest.raises(ValueError, match='Rejected Request: langfuse_host is not allowed in request') as exc: + with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", - metadata_key: json.dumps( - {"langfuse_host": "https://attacker.example"} - ), + metadata_key: json.dumps({"langfuse_host": "https://attacker.example"}), }, general_settings={}, llm_router=None, @@ -2901,7 +2806,7 @@ def test_model_level_allow_does_not_skip_subsequent_banned_params(monkeypatch): lambda model, param, request_body_value, llm_router: param == "api_base", ) - with pytest.raises(ValueError, match='Rejected Request: langfuse_host is not allowed in request') as exc: + with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -2941,8 +2846,7 @@ def test_observability_ban_covers_canonical_supported_callback_params(): ) for param in _request_blocked_callback_params: assert param in banned, ( - f"{param} is in _request_blocked_callback_params but is not banned " - "at the proxy request-body boundary." + f"{param} is in _request_blocked_callback_params but is not banned at the proxy request-body boundary." ) @@ -2972,7 +2876,7 @@ class TestPricingInjectionBlocked: ], ) def test_pricing_field_rejected_by_default(self, field, value): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: value}, general_settings={}, @@ -3040,9 +2944,7 @@ class TestGetRequestRouteTemplate: def test_exception_returns_none(self): req = MagicMock() - type(req).scope = property( - lambda self: (_ for _ in ()).throw(RuntimeError("boom")) - ) + type(req).scope = property(lambda self: (_ for _ in ()).throw(RuntimeError("boom"))) assert get_request_route_template(req) is None @@ -3095,9 +2997,7 @@ class TestGetKeyTagRateLimits: """Tests for get_key_tag_rpm_limit.""" def test_reads_tag_rpm_limit_from_metadata(self): - key = UserAPIKeyAuth( - api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}} - ) + key = UserAPIKeyAuth(api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}}) assert get_key_tag_rpm_limit(key) == {"cell-1": 5} def test_returns_none_when_unset(self): @@ -3160,12 +3060,8 @@ class TestIsRequestBodySafeChecksBracketNotationMetadata: def test_bracket_notation_matches_json_encoding_for_deeper_nesting(self): """A value nested below the first level is treated the same either way: the check descends one level into metadata, for both encodings.""" - deep_bracket = { - "litellm_metadata[spend_logs_metadata][langfuse_host]": "https://example.invalid" - } - deep_json = { - "litellm_metadata": {"spend_logs_metadata": {"langfuse_host": "https://example.invalid"}} - } + deep_bracket = {"litellm_metadata[spend_logs_metadata][langfuse_host]": "https://example.invalid"} + deep_json = {"litellm_metadata": {"spend_logs_metadata": {"langfuse_host": "https://example.invalid"}}} kwargs = dict(general_settings={}, llm_router=None, model="gpt-4") assert is_request_body_safe(request_body=deep_bracket, **kwargs) is True assert is_request_body_safe(request_body=deep_json, **kwargs) is True @@ -3214,9 +3110,7 @@ class TestHasUserSetupSso: def test_true_for_saml_metadata_url(self, monkeypatch): from litellm.proxy.auth.auth_utils import _has_user_setup_sso - monkeypatch.setenv( - "SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml" - ) + monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml") assert _has_user_setup_sso() is True def test_true_for_saml_metadata_xml(self, monkeypatch):