fix(anthropic): make the workspace id server-owned, and close two more surface gaps

Three defects the review bot found, all confirmed against the code before fixing.

anthropic_workspace_id was carved out of the request-body ban as inert. It is not
inert: it is the scope the federation token is minted for, and router.py merges
request kwargs over deployment params, so a caller who set it chose the scope
instead of the administrator. Proven against Anthropic before fixing, whose token
endpoint answered "workspace_id is not a well-formed wrkspc_ tagged ID" on a value
that came from the request body. The carve-out is gone, so it is banned like every
other minting parameter.

That does not cost the Bedrock Claude Platform route anything. It reads a workspace
from workspace_id or aws_workspace_id as well, neither of which is a federation
parameter or in all_litellm_params, so both still reach it from a request body, and
an administrator-configured anthropic_workspace_id still resolves from the
deployment. The refusal names those two spellings so that caller is not left
guessing.

Batches had the same caller-credential merge that files did, on its create path.
The retrieve path in the handler passes no caller headers, which is why the earlier
pass over this missed it.

merge_anthropic_beta_headers now accepts a list as well as a comma-separated string.
The Skills surface handled a list-valued anthropic-beta before it shared this
helper, and calling .split() on one raises.
This commit is contained in:
derhornspieler 2026-08-23 21:54:47 -04:00
parent 2bdd5becd6
commit b2fb6dd4bb
5 changed files with 186 additions and 190 deletions

View file

@ -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,

View file

@ -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))

View file

@ -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 "")
)

View file

@ -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:

View file

@ -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):