fix(anthropic): keep workload identity federation out of caller-controlled request paths

An independent review of the previous commit found the narrowing incomplete. The federation
fields select which server-side secret is read, and together with api_base decide where it is
sent, so they are deployment decisions on every surface:

- They are rejected from any request body unconditionally, ahead of the general banned-parameter
  check, because both client-side credential opt-ins would otherwise re-enable them. The inert
  workspace id keeps its existing behaviour, since Bedrock Claude Platform already accepts that
  spelling in a request body.
- /health/test_connection takes a litellm_params object that never reached the request-body check,
  and its existing guard only covers os.environ references, not oidc ones. It rejects them now.
- A client-redirected api_base clears the federation fields and marks the deployment, so a token is
  not minted for a caller-chosen host. The mark is what stops the environment-configured path,
  which cannot be cleared out of a dictionary

Batch retrieval now threads litellm_params, so it authenticates the same ways every other Anthropic
surface does rather than failing ahead of the transformation that resolves them

The exchange engine always publishes a result for its single-flight entry, so an unexpected failure
can no longer leave every later caller waiting on a leader that never finishes. Error bodies are
rendered from structured fields only, and a body that echoes the submitted assertion is dropped
rather than logged and returned. Endpoint normalization now works on the URL path, so a
single-label host is left alone and a pathological base cannot exhaust the stack
This commit is contained in:
derhornspieler 2026-08-22 22:32:25 -04:00
parent 20da15edb6
commit 79e4a6936d
10 changed files with 231 additions and 35 deletions

View file

@ -490,6 +490,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
api_key=api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
litellm_params=dict(litellm_params),
)
else:
raise litellm.exceptions.BadRequestError(

View file

@ -42,6 +42,7 @@ class AnthropicBatchesHandler:
timeout: float | httpx.Timeout,
max_retries: int | None,
logging_obj: LiteLLMLoggingObj | None = None,
litellm_params: dict | None = None,
) -> LiteLLMBatch:
"""
Async: Retrieve a batch from Anthropic.
@ -60,9 +61,7 @@ class AnthropicBatchesHandler:
# Resolve API credentials
api_base = api_base or self.anthropic_model_info.get_api_base(api_base)
api_key = api_key or self.anthropic_model_info.get_api_key()
if not api_key:
raise ValueError("Missing Anthropic API Key")
resolved_litellm_params: Final = litellm_params if litellm_params is not None else {}
# Create a minimal logging object if not provided
if logging_obj is None:
@ -85,7 +84,7 @@ class AnthropicBatchesHandler:
api_base=api_base,
batch_id=batch_id,
optional_params={},
litellm_params={},
litellm_params=resolved_litellm_params,
)
# Validate environment and get headers
@ -94,7 +93,7 @@ class AnthropicBatchesHandler:
model="",
messages=[],
optional_params={},
litellm_params={},
litellm_params=resolved_litellm_params,
api_key=api_key,
api_base=api_base,
)
@ -130,6 +129,7 @@ class AnthropicBatchesHandler:
timeout: float | httpx.Timeout,
max_retries: int | None,
logging_obj: LiteLLMLoggingObj | None = None,
litellm_params: dict | None = None,
) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]:
"""
Retrieve a batch from Anthropic.
@ -154,6 +154,7 @@ class AnthropicBatchesHandler:
timeout=timeout,
max_retries=max_retries,
logging_obj=logging_obj,
litellm_params=litellm_params,
)
else:
return asyncio.run(
@ -164,5 +165,6 @@ class AnthropicBatchesHandler:
timeout=timeout,
max_retries=max_retries,
logging_obj=logging_obj,
litellm_params=litellm_params,
)
)

View file

@ -230,14 +230,14 @@ DROP_UNSUPPORTED_SPEED_WARNING: Final = (
class AnthropicConfig(AnthropicModelInfo, BaseConfig):
_workload_identity_eligible: ClassVar[bool] = True
"""
Reference: https://docs.anthropic.com/claude/reference/messages_post
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
"""
_workload_identity_eligible: ClassVar[bool] = True
max_tokens: int | None = None
stop_sequences: list | None = None
temperature: int | None = None

View file

@ -4,6 +4,7 @@ token for a short-lived ``sk-ant-oat01`` token via the shared RFC 7523 engine.""
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, NoReturn
from urllib.parse import urlsplit, urlunsplit
from pydantic import BaseModel, ConfigDict
from typing_extensions import assert_never
@ -29,6 +30,7 @@ from litellm.types.llms.anthropic import ANTHROPIC_TOKEN_EXCHANGE_PATH
_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
_DEFAULT_API_BASE: Final = "https://api.anthropic.com"
_INLINE_ENV_VAR: Final = "ANTHROPIC_IDENTITY_TOKEN"
_DISABLE_WIF_PARAM: Final = "anthropic_disable_workload_identity_federation"
_ACCEPTED_REF_PREFIX: Final = "oidc/"
_CHAT_BASE_SUFFIXES: Final = ("/v1/messages", "/v1")
_REJECTED_REF_PREFIX: Final = "oidc/env_path/"
@ -54,6 +56,8 @@ class AnthropicWifParams(BaseModel):
def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> AnthropicWifParams | None:
if litellm_params is not None and litellm_params.get(_DISABLE_WIF_PARAM) is True:
return None
federation_rule_id: Final = _config_value(
litellm_params, "anthropic_federation_rule_id", "ANTHROPIC_FEDERATION_RULE_ID"
)
@ -151,12 +155,21 @@ def _resolve_default_api_base() -> str:
def _strip_chat_suffix(base: str) -> str:
trimmed: Final = base.rstrip("/")
stripped: Final = next(
parts: Final = urlsplit(base)
if not parts.scheme or not parts.netloc:
return base.rstrip("/")
return urlunsplit((parts.scheme, parts.netloc, _strip_path_suffixes(parts.path), "", ""))
def _strip_path_suffixes(path: str) -> str:
"""Drop the chat-surface suffixes a deployment base may carry, so every tier derives the same
token URL. Recursion depth is bounded by the path's own segment count."""
trimmed: Final = path.rstrip("/")
shortened: Final = next(
(trimmed.removesuffix(suffix) for suffix in _CHAT_BASE_SUFFIXES if trimmed.endswith(suffix)),
trimmed,
)
return trimmed if stripped == trimmed else _strip_chat_suffix(stripped)
return trimmed if shortened == trimmed else _strip_path_suffixes(shortened)
def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None:

View file

@ -57,6 +57,10 @@ _CONTENT_TYPES: Final = MappingProxyType({"json": "application/json", "form": "a
_OVERSIZED_BODY_MESSAGE: Final = "oversized error response omitted"
_NON_OBJECT_BODY_MESSAGE: Final = "non-object error response omitted"
_NO_OAUTH_FIELDS_MESSAGE: Final = "error response carried no RFC 6749 fields"
_UNSTRUCTURED_BODY_MESSAGE: Final = "non-JSON error response omitted"
_REFLECTED_VALUE_MESSAGE: Final = "<redacted: response echoed the request>"
_REFLECTION_PROBE_LENGTH: Final = 24
_SENTINEL_BODY_MESSAGES: Final = frozenset({_OVERSIZED_BODY_MESSAGE, _NON_OBJECT_BODY_MESSAGE})
class _TokenExchangeResponse(BaseModel):
@ -78,20 +82,36 @@ def validate_token_endpoint_url(url: str) -> str | InsecureTokenUrl:
return InsecureTokenUrl(host=parsed.hostname or "")
def redact_oauth_error_body(status_code: int, body_text: str) -> TokenEndpointError:
return TokenEndpointError(status_code=status_code, redacted_body=_redact_body_text(body_text))
def redact_oauth_error_body(status_code: int, body_text: str, assertion: SecretStr | None = None) -> TokenEndpointError:
rendered: Final = _redact_body_text(body_text)
return TokenEndpointError(
status_code=status_code,
redacted_body=_drop_reflected_assertion(rendered, assertion),
)
def _drop_reflected_assertion(rendered: str, assertion: SecretStr | None) -> str:
"""A token endpoint that echoes the submitted assertion back would otherwise put it in the
operator log and in the error handed to the caller."""
if assertion is None:
return rendered
secret: Final = assertion.get_secret_value()
probe: Final = secret[:_REFLECTION_PROBE_LENGTH]
if len(probe) < _REFLECTION_PROBE_LENGTH or probe not in rendered:
return rendered
return _REFLECTED_VALUE_MESSAGE
def _redact_body_text(body_text: str) -> str:
if body_text in _SENTINEL_BODY_MESSAGES:
return body_text
if len(body_text) > MAX_RESPONSE_BYTES:
return _OVERSIZED_BODY_MESSAGE
try:
parsed: Final = _REDACTABLE_BODY_ADAPTER.validate_json(body_text)
except ValidationError:
return body_text[:_REDACTION_CAP]
return _UNSTRUCTURED_BODY_MESSAGE
match parsed:
case str():
return parsed[:_REDACTION_CAP]
case Mapping():
return _format_oauth_error_fields(parsed)
case _:
@ -419,7 +439,7 @@ class JwtBearerTokenExchangeEngine:
return self._refresh_executor
def _lead(self, spec: TokenExchangeSpec, entry: _Entry) -> ExchangeResult:
result: Final = self._exchange(spec)
result: Final = self._exchange_never_raises(spec)
with self._lock:
entry.publish(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
return result
@ -439,7 +459,7 @@ class JwtBearerTokenExchangeEngine:
return TokenTransportError(detail="timed out waiting for the token exchange leader")
def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
result: Final = self._exchange(spec)
result: Final = self._exchange_never_raises(spec)
with self._lock:
entry.publish_advisory(result, backoff_until=self._clock() + ADVISORY_REFRESH_BACKOFF_SECONDS)
stale_expires_at: Final = entry.token.expires_at if entry.token is not None else None
@ -459,15 +479,31 @@ class JwtBearerTokenExchangeEngine:
ADVISORY_REFRESH_BACKOFF_SECONDS,
)
def _exchange_never_raises(self, spec: TokenExchangeSpec) -> ExchangeResult:
"""The single-flight leader and the advisory refresher must always publish a result: an
unhandled exception here would leave the entry armed (in_flight, cleared event) forever, so
every subsequent caller for this key would follow a leader that never finishes."""
try:
return self._exchange(spec)
except Exception as e: # noqa: BLE001 # a leader must resolve its entry; any failure becomes a value
return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP])
def _exchange(self, spec: TokenExchangeSpec) -> ExchangeResult:
first: Final = self._attempt_exchange(spec)
if not isinstance(first, _Unauthorized):
return first
second: Final = self._attempt_exchange(spec)
if isinstance(second, _Unauthorized):
return redact_oauth_error_body(second.response.status_code, _capped_body_text(second.response))
return redact_oauth_error_body(
second.response.status_code, _capped_body_text(second.response), self._reread_assertion(spec)
)
return second
def _reread_assertion(self, spec: TokenExchangeSpec) -> SecretStr | None:
"""Best-effort re-read, purely so a reflected assertion can be recognized in an error body."""
reread: Final = _read_assertion(self._assertion_reader, spec.assertion_ref)
return reread if isinstance(reread, SecretStr) else None
def _attempt_exchange(self, spec: TokenExchangeSpec) -> "ExchangeResult | _Unauthorized":
assertion: Final = _read_assertion(self._assertion_reader, spec.assertion_ref)
if isinstance(assertion, AssertionSourceError):
@ -486,11 +522,11 @@ class JwtBearerTokenExchangeEngine:
return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP])
if response.status_code == 401:
return _Unauthorized(response=response)
return self._parse_response(response)
return self._parse_response(response, assertion)
def _parse_response(self, response: httpx.Response) -> ExchangeResult:
def _parse_response(self, response: httpx.Response, assertion: SecretStr | None = None) -> ExchangeResult:
if not 200 <= response.status_code < 300:
return redact_oauth_error_body(response.status_code, _capped_body_text(response))
return redact_oauth_error_body(response.status_code, _capped_body_text(response), assertion)
if len(response.content) > MAX_RESPONSE_BYTES:
return MalformedTokenResponse(detail="token response body exceeds the 1 MiB cap")
try:

View file

@ -187,6 +187,23 @@ 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")
)
def reject_server_owned_wif_params(body: Mapping[str, object]) -> None:
"""Raise ``ValueError`` if a request-supplied mapping carries a server-owned workload-identity
federation field. These are never client-settable, on any surface, with or without a client-side
credential opt-in."""
for param in _ANTHROPIC_WIF_UNCONDITIONAL_BANNED:
if param in body:
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."
)
_NESTED_CONFIG_KEYS: Final[tuple[str, ...]] = ("litellm_embedding_config", "extra_body")
# Metadata containers that carry per-request configuration consumed by the
@ -317,11 +334,6 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
# so a caller-supplied value picks a transport and a callback surface the
# admin did not choose.
"rust",
# Anthropic workload-identity federation. These select which server-side secret is read
# (``anthropic_identity_token`` resolves an ``oidc/...`` reference against the proxy's own
# environment and filesystem) and, together with ``api_base``, where that secret is sent, so a
# caller-supplied value is an exfiltration primitive for any env var or mounted token file.
*sorted(anthropic_wif_litellm_params),
# SDK-only field; also rejected outright in is_request_body_safe.
"model_list",
"vertex_ai_credentials",
@ -345,6 +357,7 @@ def _check_banned_params(
Shared between the root-level check and the nested-config check so a
new banned param only needs to be added in one place.
"""
reject_server_owned_wif_params(body)
for param in _BANNED_REQUEST_BODY_PARAMS:
if param not in body:
continue
@ -485,6 +498,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
_reject_url_valued_fallback_target(target)
litellm_params: Final = _coerce_metadata_to_dict(request_body.get("litellm_params"))
if litellm_params is not None:
reject_server_owned_wif_params(litellm_params)
litellm_params_metadata: Final = _coerce_metadata_to_dict(litellm_params.get("metadata"))
if litellm_params_metadata is not None:
_check_banned_params(

View file

@ -31,6 +31,7 @@ from litellm.proxy._types import (
)
from litellm.proxy.auth.auth_utils import (
_BANNED_REQUEST_BODY_PARAMS, # pyright: ignore[reportPrivateUsage] # one canonical list, shared with the request-body check
reject_server_owned_wif_params,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
@ -1940,6 +1941,7 @@ async def test_model_connection(
"Could not find model %s in router: %s. Proceeding with request params only.", model_name, e
)
reject_server_owned_wif_params(request_litellm_params)
# Merge: config params (from proxy config) as base, request params override
litellm_params = {
**_config_base_for_health_check(

View file

@ -13,8 +13,16 @@ Ensures cooldowns are applied correctly.
from typing import Final
from litellm.types.utils import anthropic_wif_litellm_params
clientside_credential_keys: Final = ["api_key", "api_base", "base_url"]
# Set on a deployment whose api_base was client-redirected, so the Anthropic auth path refuses to
# mint a federation token there even when WIF is configured only through ANTHROPIC_* env vars (which
# cannot be cleared from litellm_params).
DISABLE_WORKLOAD_IDENTITY_PARAM: Final = "anthropic_disable_workload_identity_federation"
_ANTHROPIC_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(anthropic_wif_litellm_params))
def _admin_config_fields_to_clear_on_base_override() -> list[str]:
"""
@ -59,6 +67,11 @@ def _admin_config_fields_to_clear_on_base_override() -> list[str]:
# ``api_base`` for the same reason as the OCI entries above.
"nvcf_function_id",
"use_ssl",
# Anthropic workload-identity federation minting fields. Not declared on
# CredentialLiteLLMParams, so listed here: a federation token minted for a
# client-redirected api_base would send the workload's OIDC assertion, and
# then the minted bearer, to the caller-chosen host.
*_ANTHROPIC_WIF_CLEAR_ON_BASE_OVERRIDE,
]
return typed_fields + kwargs_only_fields
@ -101,5 +114,6 @@ def get_dynamic_litellm_params(litellm_params: dict, request_kwargs: dict) -> di
litellm_params.pop(field, None)
if field in request_kwargs:
litellm_params[field] = request_kwargs[field]
litellm_params[DISABLE_WORKLOAD_IDENTITY_PARAM] = True
return litellm_params

View file

@ -2966,10 +2966,109 @@ class TestWifExchangeTransportHardening:
class TestWifParamsAreNotClientSettable:
def test_every_wif_param_is_banned_from_request_bodies(self):
"""These fields choose which server-side secret is read and, with api_base, where it is sent,
so a caller-supplied value would be an exfiltration primitive."""
from litellm.proxy.auth.auth_utils import _BANNED_REQUEST_BODY_PARAMS
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."""
from litellm.proxy.auth.auth_utils import _ANTHROPIC_WIF_UNCONDITIONAL_BANNED
from litellm.types.utils import anthropic_wif_litellm_params
assert set(anthropic_wif_litellm_params) <= set(_BANNED_REQUEST_BODY_PARAMS)
assert set(_ANTHROPIC_WIF_UNCONDITIONAL_BANNED) == set(anthropic_wif_litellm_params) - {
"anthropic_workspace_id"
}
class TestWifServerOwnedParamsAreUnconditional:
"""The minting fields choose which server-side secret is read and, with api_base, where it goes,
so no client-side credential opt-in may re-enable them."""
@staticmethod
def _body(param: str) -> dict:
return {"model": "claude-sonnet-5", param: "oidc/env/SOME_SERVER_SECRET"}
@pytest.mark.parametrize(
"param",
[
"anthropic_identity_token",
"anthropic_identity_token_file",
"anthropic_federation_rule_id",
"anthropic_organization_id",
"anthropic_service_account_id",
],
)
def test_rejected_even_with_proxy_wide_opt_in(self, param: str):
from litellm.proxy.auth.auth_utils import is_request_body_safe
with pytest.raises(ValueError, match="server-owned workload identity federation"):
is_request_body_safe(
request_body=self._body(param),
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="claude-sonnet-5",
)
def test_rejected_inside_nested_litellm_params(self):
from litellm.proxy.auth.auth_utils import is_request_body_safe
with pytest.raises(ValueError, match="server-owned workload identity federation"):
is_request_body_safe(
request_body={"model": "claude-sonnet-5", "litellm_params": self._body("anthropic_identity_token")},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
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."""
from litellm.proxy.auth.auth_utils import is_request_body_safe
assert (
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
)
class TestWifDisabledOnClientRedirectedBase:
def test_base_override_clears_wif_and_sets_the_sentinel(self):
"""A federation token minted for a client-chosen api_base would send the workload's assertion,
and then the minted bearer, to that host."""
from litellm.llms.anthropic.wif import resolve_anthropic_wif_params
from litellm.router_utils.clientside_credential_handler import (
DISABLE_WORKLOAD_IDENTITY_PARAM,
get_dynamic_litellm_params,
)
admin_deployment = {
"model": "anthropic/claude-sonnet-5",
"anthropic_federation_rule_id": "fdrl_admin",
"anthropic_organization_id": "org-admin",
"anthropic_identity_token": "oidc/env/WIF_TEST_JWT",
}
redirected = get_dynamic_litellm_params(
litellm_params=dict(admin_deployment),
request_kwargs={"api_base": "https://not-anthropic.example"},
)
assert redirected[DISABLE_WORKLOAD_IDENTITY_PARAM] is True
assert "anthropic_federation_rule_id" not in redirected
assert resolve_anthropic_wif_params(redirected) is None
def test_sentinel_blocks_env_var_configured_federation(self, monkeypatch):
"""Environment-configured federation cannot be cleared out of a dict, so the sentinel is what
stops it on a redirected deployment."""
from litellm.llms.anthropic.wif import resolve_anthropic_wif_params
from litellm.router_utils.clientside_credential_handler import DISABLE_WORKLOAD_IDENTITY_PARAM
monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env")
monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env")
monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "oidc/env/WIF_TEST_JWT")
monkeypatch.setenv("WIF_TEST_JWT", "jwt-assertion-value")
assert resolve_anthropic_wif_params({}) is not None
assert resolve_anthropic_wif_params({DISABLE_WORKLOAD_IDENTITY_PARAM: True}) is None

View file

@ -9,6 +9,7 @@ from typing import Final
from urllib.parse import parse_qsl
import httpx
from pydantic import SecretStr
import pytest
from litellm.llms.base_llm.auth.token_exchange import (
@ -476,13 +477,27 @@ class TestRedactionAndCaps:
assert "m" * 256 in result.redacted_body
assert "m" * 257 not in result.redacted_body
def test_string_body_truncated(self):
def test_json_string_body_is_not_echoed(self):
"""A free-text body can carry back whatever was sent, so only structured OAuth fields are
ever rendered into an error an operator or caller will see."""
result = redact_oauth_error_body(400, json.dumps("s" * 500))
assert result.redacted_body == "s" * 256
assert result.redacted_body == "non-object error response omitted"
assert "s" * 32 not in result.redacted_body
def test_plain_text_body_truncated(self):
def test_plain_text_body_is_not_echoed(self):
result = redact_oauth_error_body(502, "t" * 500)
assert result.redacted_body == "t" * 256
assert result.redacted_body == "non-JSON error response omitted"
assert "t" * 32 not in result.redacted_body
def test_reflected_assertion_is_dropped(self):
"""An endpoint that echoes the submitted assertion must not put it in the log or the error."""
assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9.REFLECTEDPAYLOAD.signature")
body = {"error": "invalid_grant", "error_description": f"bad assertion {assertion.get_secret_value()}"}
result = redact_oauth_error_body(400, json.dumps(body), assertion)
assert assertion.get_secret_value() not in result.redacted_body
assert "REFLECTEDPAYLOAD" not in result.redacted_body
def test_json_array_body_constant_message(self):
result = redact_oauth_error_body(400, json.dumps(["a", "b"]))