diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index b7de263fd6c..b8c3f1917b0 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -35,6 +35,21 @@ ANTHROPIC_WIF_KWARGS_KEYS: Final = frozenset( "anthropic_workspace_id", "anthropic_identity_token_file", "anthropic_identity_token", + # Identity-source selection (Phase 1): absent means the legacy + # token_file/env resolver above, byte-identical to today. + "anthropic_identity_source", + # internal_issuer: litellm self-signs the workload assertion. + "anthropic_issuer_url", + "anthropic_issuer_subject", + "anthropic_issuer_audience", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_signing_key_ref", + # keycloak: litellm fetches the assertion via client_credentials. + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope", } ) diff --git a/litellm/llms/anthropic/wif.py b/litellm/llms/anthropic/wif.py index 837f63d3895..03340c0e084 100644 --- a/litellm/llms/anthropic/wif.py +++ b/litellm/llms/anthropic/wif.py @@ -1,15 +1,23 @@ """Anthropic workload identity federation: exchanges an external OIDC identity token for a short-lived ``sk-ant-oat01`` token via the shared RFC 7523 engine.""" -from collections.abc import Mapping +from collections.abc import Callable, Mapping from types import MappingProxyType -from typing import Final, NoReturn +from typing import Final, NoReturn, TypeVar from urllib.parse import urlsplit, urlunsplit -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, ValidationError from typing_extensions import assert_never import litellm +from litellm.llms.base_llm.auth.client_credentials import keycloak_assertion_source +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_ref, +) +from litellm.llms.base_llm.auth.internal_issuer import internal_issuer_assertion_source from litellm.llms.base_llm.auth.token_exchange import ( JwtBearerTokenExchangeEngine, default_token_exchange_engine, @@ -34,6 +42,31 @@ _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/" +_IDENTITY_SOURCE_PARAM: Final = "anthropic_identity_source" +_IDENTITY_SOURCE_ENV: Final = "ANTHROPIC_IDENTITY_SOURCE" + +# litellm_params key -> InternalIssuerSource/KeycloakSource field name. Every key here must +# also be listed in ANTHROPIC_WIF_KWARGS_KEYS (get_litellm_params.py), which is what makes it +# request-banned and cleared on a client-redirected api_base -- see types/utils.py's +# anthropic_wif_litellm_params, derived from that same set. +_INTERNAL_ISSUER_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType( + { + "anthropic_issuer_url": "issuer_url", + "anthropic_issuer_subject": "subject", + "anthropic_issuer_audience": "audience", + "anthropic_issuer_ttl_seconds": "ttl_seconds", + "anthropic_issuer_signing_key_ref": "signing_key_ref", + } +) +_KEYCLOAK_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType( + { + "anthropic_keycloak_token_url": "token_url", + "anthropic_keycloak_client_id": "client_id", + "anthropic_keycloak_auth_method": "auth_method", + "anthropic_keycloak_client_secret_ref": "client_secret_ref", + "anthropic_keycloak_scope": "scope", + } +) _WORKSPACE_HINT: Final = ( " If the federation rule is scoped to a workspace, set ANTHROPIC_WORKSPACE_ID" " (or the anthropic_workspace_id litellm param) to that workspace id." @@ -43,6 +76,9 @@ _ALLOWLIST_HINT: Final = ( " (/var/run/secrets or /run/secrets by default);" " set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist." ) +_EMPTY_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) + +_IdentitySourceVariant = TypeVar("_IdentitySourceVariant", bound="InternalIssuerSource | KeycloakSource") class AnthropicWifParams(BaseModel): @@ -53,6 +89,7 @@ class AnthropicWifParams(BaseModel): service_account_id: str | None = None workspace_id: str | None = None assertion_ref: str + assertion_source: Callable[[], str | None] | None = None def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> AnthropicWifParams | None: @@ -64,9 +101,10 @@ def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID") if federation_rule_id is None or organization_id is None: return None - assertion_ref: Final = _resolve_assertion_ref(litellm_params) - if assertion_ref is None: + identity_source: Final = _resolve_identity_source(litellm_params) + if identity_source is None: return None + assertion_ref, assertion_source = identity_source return AnthropicWifParams( federation_rule_id=federation_rule_id, organization_id=organization_id, @@ -75,9 +113,82 @@ def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> ), workspace_id=_config_value(litellm_params, "anthropic_workspace_id", "ANTHROPIC_WORKSPACE_ID"), assertion_ref=assertion_ref, + assertion_source=assertion_source, ) +def _resolve_identity_source( + litellm_params: Mapping[str, object] | None, +) -> tuple[str, Callable[[], str] | None] | None: + """Dispatches on ``anthropic_identity_source``. Absent (the default) keeps today's + token_file/env resolution byte-identical, with no ``assertion_source`` closure -- the engine + falls back to its own reader exactly as it does today. A recognized kind builds the matching + frozen config, hashes it into the ``oidc//`` cache-key ref (``identity_source_ref``), + and closes the source's fetch/mint function over it. An unset-but-invalid config (unknown + kind, a missing required field, or a field from the other variant) fails closed here rather + than silently falling back to token_file.""" + source_kind: Final = _config_value(litellm_params, _IDENTITY_SOURCE_PARAM, _IDENTITY_SOURCE_ENV) + if source_kind is None: + legacy_ref: Final = _resolve_assertion_ref(litellm_params) + return (legacy_ref, None) if legacy_ref is not None else None + params: Final = litellm_params if litellm_params is not None else _EMPTY_PARAMS + match source_kind: + case AnthropicIdentitySourceKind.internal_issuer.value: + _reject_foreign_variant_fields(params, foreign_field_map=_KEYCLOAK_FIELD_MAP, chosen_kind=source_kind) + issuer_config: Final = _build_variant(InternalIssuerSource, params, _INTERNAL_ISSUER_FIELD_MAP) + return identity_source_ref(issuer_config), internal_issuer_assertion_source(issuer_config) + case AnthropicIdentitySourceKind.keycloak.value: + _reject_foreign_variant_fields( + params, foreign_field_map=_INTERNAL_ISSUER_FIELD_MAP, chosen_kind=source_kind + ) + keycloak_config: Final = _build_variant(KeycloakSource, params, _KEYCLOAK_FIELD_MAP) + return identity_source_ref(keycloak_config), keycloak_assertion_source(keycloak_config) + case _: + raise litellm.AuthenticationError( + message=( + f"{_IDENTITY_SOURCE_PARAM} must be one of " + f"{', '.join(kind.value for kind in AnthropicIdentitySourceKind)}; got {source_kind!r}." + ), + llm_provider="anthropic", + model="", + ) + + +def _reject_foreign_variant_fields( + litellm_params: Mapping[str, object], foreign_field_map: Mapping[str, str], chosen_kind: str +) -> None: + foreign_keys_present: Final = tuple(param for param in foreign_field_map if param in litellm_params) + if foreign_keys_present: + raise litellm.AuthenticationError( + message=( + f"{_IDENTITY_SOURCE_PARAM} is {chosen_kind!r}, but {', '.join(sorted(foreign_keys_present))} " + "belongs to a different identity source and cannot be set alongside it." + ), + llm_provider="anthropic", + model="", + ) + + +def _build_variant( + model: type[_IdentitySourceVariant], + litellm_params: Mapping[str, object], + field_map: Mapping[str, str], +) -> _IdentitySourceVariant: + fields: Final = MappingProxyType( + {field_map[key]: value for key, value in litellm_params.items() if key in field_map} + ) + try: + return model.model_validate(fields) + except ValidationError as e: + # hide_input_in_errors=True on both variant models keeps a secret pasted into the + # wrong field (e.g. a client_secret typed as signing_key_ref) out of str(e). + raise litellm.AuthenticationError( + message=f"Invalid {_IDENTITY_SOURCE_PARAM} configuration: {e}", + llm_provider="anthropic", + model="", + ) from e + + def build_anthropic_wif_spec(params: AnthropicWifParams, api_base: str) -> TokenExchangeSpec: return TokenExchangeSpec( token_url=api_base.rstrip("/") + ANTHROPIC_TOKEN_EXCHANGE_PATH, @@ -98,6 +209,7 @@ def build_anthropic_wif_spec(params: AnthropicWifParams, api_base: str) -> Token ), body_encoding="json", request_headers=MappingProxyType({}), + assertion_source=params.assertion_source, cache_key_identity=( params.federation_rule_id, params.organization_id, @@ -233,7 +345,8 @@ def _error_detail(error: ExchangeError, workspace_id_set: bool) -> str: case AssertionSourceError() if error.kind == "disallowed_path": return f"Could not read the OIDC identity token from {error.source_ref}.{_ALLOWLIST_HINT}" case AssertionSourceError(): - return f"Could not obtain the OIDC identity token ({error.kind}) from {error.source_ref}." + base: Final = f"Could not obtain the OIDC identity token ({error.kind}) from {error.source_ref}." + return f"{base} {error.detail}" if error.detail else base case InsecureTokenUrl(): return f"The token endpoint must use https; refusing to send the identity token to host {error.host!r}." case TokenEndpointError() if error.status_code == 401 and not workspace_id_set: diff --git a/litellm/llms/base_llm/auth/__init__.py b/litellm/llms/base_llm/auth/__init__.py index 985e093b2bb..291e1492c98 100644 --- a/litellm/llms/base_llm/auth/__init__.py +++ b/litellm/llms/base_llm/auth/__init__.py @@ -1,3 +1,31 @@ +from litellm.llms.base_llm.auth.client_credentials import ( + SecretReader, + fetch_keycloak_assertion, + keycloak_assertion_source, +) +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceConfig, + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_config_adapter, + identity_source_ref, +) +from litellm.llms.base_llm.auth.internal_issuer import ( + SigningKeyReader, + internal_issuer_assertion_source, + internal_issuer_jwks_document, + mint_internal_issuer_assertion, +) +from litellm.llms.base_llm.auth.jwt_signing import ( + ALG, + build_jwk, + build_jwks, + jwks_document_json, + load_es256_private_key, + rfc7638_thumbprint, + sign_es256_jwt, +) from litellm.llms.base_llm.auth.token_exchange import ( ADVISORY_REFRESH_BACKOFF_SECONDS, ADVISORY_REFRESH_SECONDS, @@ -11,6 +39,7 @@ from litellm.llms.base_llm.auth.token_exchange import ( ) from litellm.llms.base_llm.auth.types import ( AssertionReader, + AssertionSource, AssertionSourceError, BodyEncoding, ExchangeError, @@ -27,23 +56,44 @@ from litellm.llms.base_llm.auth.types import ( __all__ = ( "ADVISORY_REFRESH_BACKOFF_SECONDS", "ADVISORY_REFRESH_SECONDS", + "ALG", "MANDATORY_REFRESH_SECONDS", "MAX_ASSERTION_BYTES", "MAX_RESPONSE_BYTES", + "AnthropicIdentitySourceConfig", + "AnthropicIdentitySourceKind", "AssertionReader", + "AssertionSource", "AssertionSourceError", "BodyEncoding", "ExchangeError", "ExchangeResult", "InsecureTokenUrl", + "InternalIssuerSource", "JwtBearerTokenExchangeEngine", + "KeycloakSource", "MalformedTokenResponse", "MintedToken", + "SecretReader", + "SigningKeyReader", "SyncTokenPoster", "TokenEndpointError", "TokenExchangeSpec", "TokenTransportError", + "build_jwk", + "build_jwks", "default_token_exchange_engine", + "fetch_keycloak_assertion", + "identity_source_config_adapter", + "identity_source_ref", + "internal_issuer_assertion_source", + "internal_issuer_jwks_document", + "jwks_document_json", + "keycloak_assertion_source", + "load_es256_private_key", + "mint_internal_issuer_assertion", "redact_oauth_error_body", + "rfc7638_thumbprint", + "sign_es256_jwt", "validate_token_endpoint_url", ) diff --git a/litellm/llms/base_llm/auth/client_credentials.py b/litellm/llms/base_llm/auth/client_credentials.py new file mode 100644 index 00000000000..1b70a0a1253 --- /dev/null +++ b/litellm/llms/base_llm/auth/client_credentials.py @@ -0,0 +1,187 @@ +"""Fetches a fresh RFC 6749 client_credentials assertion for Anthropic's ``keycloak`` identity +source: LiteLLM authenticates to Keycloak as its own confidential client and presents the +resulting ``access_token`` as the workload assertion (Phase 1 decision 2). + +The client secret is the operator-supplied pointer at ``KeycloakSource.client_secret_ref``, +resolved the same way every other WIF secret pointer already is (env, a Credential, or whatever +secret manager ``litellm.secret_manager_client`` is globally configured to, Vault included). +Every fetch is a fresh HTTP POST; nothing here caches a fetched token, since the outer +token-exchange engine already caches the Anthropic token it buys with one -- see decision 2's +"no Keycloak-side cache" ruling. +""" + +import base64 +import threading +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, TypeAlias +from urllib.parse import urlencode + +import httpx +from pydantic import BaseModel, SecretStr, ValidationError +from typing_extensions import assert_never + +from litellm.llms.base_llm.auth.identity_source import KeycloakSource +from litellm.llms.base_llm.auth.token_exchange import ( + MAX_RESPONSE_BYTES, + redact_oauth_error_body, + validate_token_endpoint_url, +) +from litellm.llms.base_llm.auth.types import InsecureTokenUrl, SyncTokenPoster + +if TYPE_CHECKING: + from litellm.llms.custom_httpx.http_handler import HTTPHandler + +SecretReader: TypeAlias = Callable[[str], str | None] # mutable-ok: Callable param-list syntax, not a list + +_GRANT_TYPE: Final = "client_credentials" +_TIMEOUT_SECONDS: Final = 30.0 +_FORM_CONTENT_TYPE: Final = "application/x-www-form-urlencoded" + + +class _ClientCredentialsResponse(BaseModel): + access_token: str + + +def _default_secret_reader(ref: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + return get_secret_str(ref) + + +class _HttpxSyncKeycloakPoster: + """Dedicated HTTPHandler for the Keycloak token POST: no ``logging_obj`` (so litellm's + request/response logging never sees the client secret or the fetched token), redirects + disabled. A separate instance from the outer engine's own poster, since this is a genuinely + new HTTP call site whose no-logging guarantee must be built here, not assumed inherited.""" + + def __init__(self) -> None: + self._lock: Final = threading.Lock() + self._handler: HTTPHandler | None = None + + def _handler_instance(self) -> "HTTPHandler": + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + with self._lock: + if self._handler is None: + handler: Final = HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0)) + handler.client.follow_redirects = False + self._handler = handler + return self._handler + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + try: + response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below + url, + content=content, + headers=dict(headers), # mutable-ok: HTTPHandler.post requires a concrete dict + timeout=timeout, + ) + except httpx.HTTPStatusError as e: + return e.response + if response is None: + raise httpx.TransportError("keycloak token endpoint returned no response") + return response + + +_DEFAULT_POSTER: Final[SyncTokenPoster] = _HttpxSyncKeycloakPoster() + + +def _basic_auth_header(client_id: str, client_secret: str) -> str: + return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode("ascii") + + +def _prepared_request(config: KeycloakSource, client_secret: str) -> tuple[bytes, Mapping[str, str]]: + scope_field: Final[Mapping[str, str]] = ( + MappingProxyType({"scope": config.scope}) if config.scope else MappingProxyType({}) + ) + match config.auth_method: + case "client_secret_basic": + return ( + urlencode(MappingProxyType({"grant_type": _GRANT_TYPE, **scope_field})).encode(), + MappingProxyType( + { + "content-type": _FORM_CONTENT_TYPE, + "authorization": _basic_auth_header(config.client_id, client_secret), + } + ), + ) + case "client_secret_post": + return ( + urlencode( + MappingProxyType( + { + "grant_type": _GRANT_TYPE, + "client_id": config.client_id, + "client_secret": client_secret, + **scope_field, + } + ) + ).encode(), + MappingProxyType({"content-type": _FORM_CONTENT_TYPE}), + ) + case _: + assert_never(config.auth_method) + + +def _resolve_client_secret(config: KeycloakSource, secret_reader: SecretReader) -> str: + secret: Final = secret_reader(config.client_secret_ref) + if not secret: + raise ValueError(f"keycloak client secret {config.client_secret_ref} could not be read") + return secret + + +def _endpoint_error_message(config: KeycloakSource, response: httpx.Response, client_secret: str) -> str: + endpoint_error: Final = redact_oauth_error_body(response.status_code, response.text, SecretStr(client_secret)) + return f"keycloak token endpoint {config.token_url} returned HTTP {endpoint_error.status_code}: {endpoint_error.redacted_body}" + + +def _parse_success_body(response: httpx.Response) -> str: + if len(response.content) > MAX_RESPONSE_BYTES: + raise ValueError("keycloak token response exceeded the size cap") + try: + parsed: Final = _ClientCredentialsResponse.model_validate_json(response.content) + except ValidationError as e: + raise ValueError("keycloak token response failed schema validation") from e + token: Final = parsed.access_token.strip() + if not token: + raise ValueError("keycloak token response carried an empty access_token") + return token + + +def fetch_keycloak_assertion( + config: KeycloakSource, + *, + poster: SyncTokenPoster = _DEFAULT_POSTER, + secret_reader: SecretReader = _default_secret_reader, +) -> str: + """POSTs one fresh client_credentials grant and returns the resulting ``access_token`` as the + workload assertion; the caller must not cache the result -- see the module docstring.""" + match validate_token_endpoint_url(config.token_url): + case InsecureTokenUrl(host=host): + raise ValueError( + f"keycloak token_url must use https; refusing to send the client secret to host {host!r}" + ) + case _: + pass + client_secret: Final = _resolve_client_secret(config, secret_reader) + content, headers = _prepared_request(config, client_secret) + try: + response: Final = poster.post(config.token_url, content=content, headers=headers, timeout=_TIMEOUT_SECONDS) + except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; every failure becomes a ValueError + raise ValueError(f"could not reach the keycloak token endpoint {config.token_url}: {type(e).__name__}") from e + if not 200 <= response.status_code < 300: + raise ValueError(_endpoint_error_message(config, response, client_secret)) + return _parse_success_body(response) + + +def keycloak_assertion_source( + config: KeycloakSource, + *, + poster: SyncTokenPoster = _DEFAULT_POSTER, + secret_reader: SecretReader = _default_secret_reader, +) -> Callable[[], str]: + """A zero-arg closure that fetches fresh on every call: the shape an ``oidc/keycloak/...`` + ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7) + -- the caller parses the config and closes this function over it, with no registry involved.""" + return lambda: fetch_keycloak_assertion(config, poster=poster, secret_reader=secret_reader) diff --git a/litellm/llms/base_llm/auth/identity_source.py b/litellm/llms/base_llm/auth/identity_source.py new file mode 100644 index 00000000000..9b9df7783f9 --- /dev/null +++ b/litellm/llms/base_llm/auth/identity_source.py @@ -0,0 +1,56 @@ +"""Tagged-union identity-source configs for Anthropic workload identity federation, beyond the +existing token_file/env resolver in ``litellm/llms/anthropic/wif.py``. + +Each variant only ever carries secret *pointer names* (``signing_key_ref``, ``client_secret_ref``), +never a resolved secret value, so ``identity_source_ref`` can safely hash a variant into the short, +content-derived ``oidc//`` string used elsewhere as a get_secret ref, a token-exchange +cache-key discriminator, and an operator-facing error pointer. +""" + +import hashlib +from enum import Enum +from typing import Annotated, Final, Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter + +_REF_HASH_HEX_LENGTH: Final = 16 +_MAX_TTL_SECONDS: Final = 3600 +_DEFAULT_TTL_SECONDS: Final = 300 + + +class AnthropicIdentitySourceKind(str, Enum): + internal_issuer = "internal_issuer" + keycloak = "keycloak" + + +class InternalIssuerSource(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True) + + kind: Literal[AnthropicIdentitySourceKind.internal_issuer] = AnthropicIdentitySourceKind.internal_issuer + issuer_url: str + subject: str + audience: str | None = None + ttl_seconds: Annotated[int, Field(gt=0, le=_MAX_TTL_SECONDS)] = _DEFAULT_TTL_SECONDS + signing_key_ref: str + + +class KeycloakSource(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True) + + kind: Literal[AnthropicIdentitySourceKind.keycloak] = AnthropicIdentitySourceKind.keycloak + token_url: str + client_id: str + auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic" + client_secret_ref: str + scope: str | None = None + + +AnthropicIdentitySourceConfig: TypeAlias = Annotated[InternalIssuerSource | KeycloakSource, Field(discriminator="kind")] +identity_source_config_adapter: Final = TypeAdapter[AnthropicIdentitySourceConfig](AnthropicIdentitySourceConfig) + + +def identity_source_ref(config: AnthropicIdentitySourceConfig) -> str: + """``oidc//``: a short, secret-free pointer, stable for identical config and rolling + whenever any field does, including a ``*_ref`` pointer NAME (never the secret it points to).""" + digest: Final = hashlib.sha256(config.model_dump_json().encode()).hexdigest()[:_REF_HASH_HEX_LENGTH] + return f"oidc/{config.kind.value}/{digest}" diff --git a/litellm/llms/base_llm/auth/internal_issuer.py b/litellm/llms/base_llm/auth/internal_issuer.py new file mode 100644 index 00000000000..6b350ef6330 --- /dev/null +++ b/litellm/llms/base_llm/auth/internal_issuer.py @@ -0,0 +1,84 @@ +"""Mints a self-issued workload assertion for Anthropic's ``internal_issuer`` identity source: +LiteLLM signs its own short-lived ES256 JWT instead of reading one from a mounted OIDC file. + +Signing custody is the operator-supplied PEM at ``InternalIssuerSource.signing_key_ref``, +resolved the same way every other WIF secret pointer already is (env, a Credential, or +whatever secret manager ``litellm.secret_manager_client`` is globally configured to, Vault +included) -- see Phase 1 decision 1. Every mint is fresh; nothing here caches a minted JWT, +since the outer token-exchange engine already caches the Anthropic token it buys with one. +""" + +import time +import uuid +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import Final, TypeAlias + +from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource +from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json, sign_es256_jwt + +SigningKeyReader: TypeAlias = Callable[[str], str | None] # mutable-ok: Callable param-list syntax, not a list + + +def _default_signing_key_reader(ref: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + return get_secret_str(ref) + + +def _claims(config: InternalIssuerSource, issued_at: int) -> Mapping[str, object]: + return MappingProxyType( + { + key: value + for key, value in ( + ("sub", config.subject), + ("iss", config.issuer_url), + ("aud", config.audience), + ("iat", issued_at), + ("exp", issued_at + config.ttl_seconds), + ("jti", str(uuid.uuid4())), + ) + if value is not None + } + ) + + +def _resolve_signing_key(config: InternalIssuerSource, key_reader: SigningKeyReader) -> str: + pem: Final = key_reader(config.signing_key_ref) + if not pem: + raise ValueError(f"internal_issuer signing key {config.signing_key_ref} could not be read") + return pem + + +def mint_internal_issuer_assertion( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, + clock: Callable[[], float] = time.time, +) -> str: + """Signs one fresh, short-lived assertion; the caller must not cache the result, since a + cached copy would defeat the point of re-minting on every exchange.""" + pem: Final = _resolve_signing_key(config, key_reader) + return sign_es256_jwt(pem, _claims(config, issued_at=int(clock()))) + + +def internal_issuer_assertion_source( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, + clock: Callable[[], float] = time.time, +) -> Callable[[], str]: + """A zero-arg closure that mints fresh on every call: the shape an ``oidc/internal_issuer/...`` + ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7) + -- the caller parses the config and closes this function over it, with no registry involved.""" + return lambda: mint_internal_issuer_assertion(config, key_reader=key_reader, clock=clock) + + +def internal_issuer_jwks_document( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, +) -> str: + """The operator-facing JWKS export, resolved from a configured identity source rather than + a raw PEM in hand -- the JSON document to register as Anthropic's inline federation issuer.""" + return jwks_document_json(_resolve_signing_key(config, key_reader)) diff --git a/litellm/llms/base_llm/auth/jwt_signing.py b/litellm/llms/base_llm/auth/jwt_signing.py new file mode 100644 index 00000000000..a7e7f8272f9 --- /dev/null +++ b/litellm/llms/base_llm/auth/jwt_signing.py @@ -0,0 +1,103 @@ +"""ES256 JWT signing primitives for Anthropic workload identity federation's +``internal_issuer`` identity source (see ``identity_source.InternalIssuerSource``). + +Pure functions over an already-resolved PEM string: no I/O, no secret-manager awareness, no +caching. Given the signing key at, say, $ISSUER_SIGNING_KEY_PEM, an operator publishes the +JWKS document Anthropic's inline federation issuer needs with one line: + + python -c "from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json; \\ + import os; print(jwks_document_json(os.environ['ISSUER_SIGNING_KEY_PEM']))" +""" + +import base64 +import hashlib +import json +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final, TypeAlias + +import jwt +from cryptography.hazmat.primitives.asymmetric import ec +from cryptography.hazmat.primitives.serialization import load_pem_private_key + +ALG: Final = "ES256" +_JWK_CURVE_NAME: Final = "P-256" +_JWK_KEY_TYPE: Final = "EC" +_COORDINATE_BYTE_LENGTH: Final = 32 # P-256 field element width, RFC 7518 6.2.1.2/6.2.1.3 + +Jwk: TypeAlias = Mapping[str, str] +Jwks: TypeAlias = Mapping[str, tuple[Jwk, ...]] + + +def load_es256_private_key(pem: str) -> ec.EllipticCurvePrivateKey: + """Parses an unencrypted PEM EC private key. Never echoes the key material in an error.""" + try: + key: Final = load_pem_private_key(pem.encode(), password=None) + except (ValueError, TypeError) as e: + raise ValueError("internal_issuer signing key is not a valid unencrypted PEM private key") from e + if not isinstance(key, ec.EllipticCurvePrivateKey) or not isinstance(key.curve, ec.SECP256R1): + raise ValueError( # noqa: TRY004 # the reader classifies ValueError into a readable config error; TypeError would not + "internal_issuer signing key must be an EC P-256 (secp256r1) private key for ES256" + ) + return key + + +def _b64url_coordinate(value: int) -> str: + return base64.urlsafe_b64encode(value.to_bytes(_COORDINATE_BYTE_LENGTH, "big")).rstrip(b"=").decode("ascii") + + +def _jwk_thumbprint_members(public_key: ec.EllipticCurvePublicKey) -> Jwk: + """RFC 7638 3.2's exact EC member set (crv, kty, x, y) and nothing else: an extra member + here would change the thumbprint and desync it from the ``kid`` published in the JWKS.""" + numbers: Final = public_key.public_numbers() + return MappingProxyType( + { + "crv": _JWK_CURVE_NAME, + "kty": _JWK_KEY_TYPE, + "x": _b64url_coordinate(numbers.x), + "y": _b64url_coordinate(numbers.y), + } + ) + + +def rfc7638_thumbprint(public_key: ec.EllipticCurvePublicKey) -> str: + """RFC 7638: SHA-256 over the lexicographically member-ordered, whitespace-free JSON + rendering of the thumbprint members, base64url-encoded without padding.""" + canonical: Final = json.dumps( + dict(sorted(_jwk_thumbprint_members(public_key).items())), # mutable-ok: json.dumps needs a real dict + separators=(",", ":"), + ) + return base64.urlsafe_b64encode(hashlib.sha256(canonical.encode()).digest()).rstrip(b"=").decode("ascii") + + +def build_jwk(public_key: ec.EllipticCurvePublicKey, kid: str) -> Jwk: + return MappingProxyType({**_jwk_thumbprint_members(public_key), "use": "sig", "alg": ALG, "kid": kid}) + + +def build_jwks(public_key: ec.EllipticCurvePublicKey) -> Jwks: + kid: Final = rfc7638_thumbprint(public_key) + return MappingProxyType({"keys": (build_jwk(public_key, kid),)}) + + +def jwks_document_json(pem: str) -> str: + """The operator-facing export: the JSON document to register as Anthropic's inline JWKS. + + ``build_jwks`` returns ``MappingProxyType``/tuple values per this repo's no-mutation + convention; the ``json`` module only knows plain ``dict``/``list``, so those are converted + at this one serialization boundary rather than giving up immutability throughout the module. + """ + key: Final = load_es256_private_key(pem) + jwks: Final = build_jwks(key.public_key()) + return json.dumps( + {"keys": [dict(jwk) for jwk in jwks["keys"]]}, # mutable-ok: json.dumps needs real dicts/lists + indent=2, + ) + + +def sign_es256_jwt(pem: str, claims: Mapping[str, object]) -> str: + """Signs ``claims`` with the PEM key, stamping ``kid`` as its RFC 7638 thumbprint so a + verifier can look the signing key up in the published JWKS by ``kid`` alone.""" + key: Final = load_es256_private_key(pem) + kid: Final = rfc7638_thumbprint(key.public_key()) + headers: Final = {"kid": kid} # mutable-ok: PyJWT requires a real dict, not a Mapping + return jwt.encode(dict(claims), key, algorithm=ALG, headers=headers) # mutable-ok: PyJWT requires a real dict diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py index d6933d131f8..68268e6b825 100644 --- a/litellm/llms/base_llm/auth/token_exchange.py +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -26,6 +26,7 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.llms.base_llm.auth.types import ( AssertionReader, + AssertionSource, AssertionSourceError, ExchangeError, ExchangeResult, @@ -92,12 +93,16 @@ def redact_oauth_error_body(status_code: int, body_text: str, assertion: SecretS 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.""" + operator log and in the error handed to the caller. The probe caps at + ``_REFLECTION_PROBE_LENGTH`` for a long assertion (a JWT), but a short reflected value (e.g. a + Keycloak client_secret under that length) still needs the whole thing checked -- capping the + minimum here too would silently stop redacting exactly the short secrets most likely to be + hand-set rather than generated.""" 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: + if not probe or probe not in rendered: return rendered return _REFLECTED_VALUE_MESSAGE @@ -165,15 +170,23 @@ def _cache_key(spec: TokenExchangeSpec) -> str: ).hexdigest() -def _read_assertion(reader: AssertionReader, ref: str) -> SecretStr | AssertionSourceError: +def _assertion_fetch(reader: AssertionReader, spec: TokenExchangeSpec) -> AssertionSource: + """``spec.assertion_source`` (an identity source's own fetch/mint closure) takes priority over + the engine-level reader when set; either way, failures are reported against ``spec.assertion_ref``.""" + if spec.assertion_source is not None: + return spec.assertion_source + return lambda: reader(spec.assertion_ref) + + +def _read_assertion(fetch: AssertionSource, ref: str) -> SecretStr | AssertionSourceError: from litellm.secret_managers.main import OidcPathNotAllowedError try: - raw: Final = reader(ref) + raw: Final = fetch() except OidcPathNotAllowedError: return AssertionSourceError(kind="disallowed_path", source_ref=ref) - except ValueError: - return AssertionSourceError(kind="unreadable", source_ref=ref) + except ValueError as e: + return AssertionSourceError(kind="unreadable", source_ref=ref, detail=str(e)[:_REDACTION_CAP]) except Exception: # noqa: BLE001 # injected readers (secret managers) raise arbitrarily; all failures become values return AssertionSourceError(kind="unreadable", source_ref=ref) if raw is None: @@ -501,11 +514,11 @@ class JwtBearerTokenExchangeEngine: 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) + reread: Final = _read_assertion(_assertion_fetch(self._assertion_reader, spec), 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) + assertion: Final = _read_assertion(_assertion_fetch(self._assertion_reader, spec), spec.assertion_ref) if isinstance(assertion, AssertionSourceError): return assertion url_check: Final = validate_token_endpoint_url(spec.token_url) diff --git a/litellm/llms/base_llm/auth/types.py b/litellm/llms/base_llm/auth/types.py index 49183f2fd8f..21a2801bdd4 100644 --- a/litellm/llms/base_llm/auth/types.py +++ b/litellm/llms/base_llm/auth/types.py @@ -8,6 +8,8 @@ import httpx from pydantic import SecretStr BodyEncoding: TypeAlias = Literal["json", "form"] +AssertionReader: TypeAlias = Callable[[str], str | None] # mutable-ok: Callable param-list syntax, not a list +AssertionSource: TypeAlias = Callable[[], str | None] # mutable-ok: Callable param-list syntax, not a list @dataclass(frozen=True, slots=True) @@ -16,6 +18,12 @@ class TokenExchangeSpec: ``token_url`` must be derived from deployment config/env only, never per-request caller input. ``assertion_ref`` is a ``oidc/...`` get_secret ref resolved fresh on every exchange. + + ``assertion_source``, when set, is a zero-arg per-config fetch/mint closure that the engine + prefers over its own engine-level ``AssertionReader`` -- the dispatch mechanism identity + sources beyond token_file/env (e.g. ``internal_issuer``, ``keycloak``) use to plug into the + shared engine without a global registry. ``assertion_ref`` still names the cache-key + discriminator and the ref echoed into operator-facing errors either way. """ token_url: str @@ -26,6 +34,7 @@ class TokenExchangeSpec: request_headers: Mapping[str, str] cache_key_identity: tuple[str, ...] timeout_seconds: float = 30.0 + assertion_source: AssertionSource | None = None @dataclass(frozen=True, slots=True) @@ -38,6 +47,7 @@ class MintedToken: class AssertionSourceError: kind: Literal["missing", "empty", "oversized", "unreadable", "disallowed_path"] source_ref: str + detail: str | None = None @dataclass(frozen=True, slots=True) @@ -71,6 +81,3 @@ class SyncTokenPoster(Protocol): """Returns the response for ANY status; never raises for status.""" def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: ... - - -AssertionReader: TypeAlias = Callable[[str], str | None] # mutable-ok: Callable param-list syntax, not a list diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index 66d7da668d4..8c7c2fabdd1 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -67,10 +67,13 @@ 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 workload-identity federation minting fields, restated here from + # anthropic_wif_litellm_params the same way azure_ad_token above is restated + # despite also being declared on CredentialLiteLLMParams (hence covered by + # typed_fields too): 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, so this list must stay correct even if a field is ever + # dropped from the typed model. *_ANTHROPIC_WIF_CLEAR_ON_BASE_OVERRIDE, ] return typed_fields + kwargs_only_fields diff --git a/litellm/types/router.py b/litellm/types/router.py index 99a4603ae49..5f331009754 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -272,6 +272,28 @@ class CredentialLiteLLMParams(BaseModel): ## IBM WATSONX ## watsonx_region_name: str | None = None + ## ANTHROPIC WORKLOAD IDENTITY FEDERATION ## + # Without these, get_deployment_credentials_with_provider silently drops a + # litellm_params-configured WIF setup before files/batches/passthrough callers see + # it, the same #30235-shaped gap azure_ad_token above was added to close. + anthropic_federation_rule_id: str | None = None + anthropic_organization_id: str | None = None + anthropic_service_account_id: str | None = None + anthropic_workspace_id: str | None = None + anthropic_identity_token_file: str | None = None + anthropic_identity_token: str | None = None + anthropic_identity_source: str | None = None + anthropic_issuer_url: str | None = None + anthropic_issuer_subject: str | None = None + anthropic_issuer_audience: str | None = None + anthropic_issuer_ttl_seconds: int | None = None + anthropic_issuer_signing_key_ref: str | None = None + anthropic_keycloak_token_url: str | None = None + anthropic_keycloak_client_id: str | None = None + anthropic_keycloak_auth_method: str | None = None + anthropic_keycloak_client_secret_ref: str | None = None + anthropic_keycloak_scope: str | None = None + _RESERVED_INIT_KEYS: Final = frozenset({"self", "params", "__class__"}) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 695d671d487..a456fcd3297 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3479,15 +3479,18 @@ bedrock_batch_litellm_params: Final = ( # Anthropic workload identity federation config, read from litellm_params by the # Anthropic auth tier. Listed for the same reason as the fields above: an # unrecognized top-level key is swept into extra_body and sent to /v1/messages. -anthropic_wif_litellm_params: Final = ( - "anthropic_federation_rule_id", - "anthropic_organization_id", - "anthropic_service_account_id", - "anthropic_workspace_id", - "anthropic_identity_token_file", - "anthropic_identity_token", +# Derived from get_litellm_params.ANTHROPIC_WIF_KWARGS_KEYS (not hand-typed) so the +# request-body ban list and the clear-on-api_base-override list can never drift from +# the set the kwargs funnel actually forwards. Imported here rather than at module top: +# get_litellm_params.py's own import chain (llms/openai/data_residency -> llms/__init__) +# reaches back into this module for CallTypes, which by this point in the file is +# already bound on the partially-initialized module. +from ..litellm_core_utils.get_litellm_params import ( # noqa: E402 # deferred past CallTypes to break the import cycle + ANTHROPIC_WIF_KWARGS_KEYS, ) +anthropic_wif_litellm_params: Final = tuple(sorted(ANTHROPIC_WIF_KWARGS_KEYS)) + all_litellm_params = ( agentic_loop_internal_litellm_params + [TRUSTED_CALLBACK_VARS_FIELD, *bedrock_batch_litellm_params, *anthropic_wif_litellm_params] diff --git a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py index a1c323b977a..3ccb2632fb4 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py @@ -281,3 +281,52 @@ class TestAnthropicWifKeys: params = get_litellm_params() for key in self.SIX_KEYS: assert key not in params + + +class TestAnthropicWifIdentitySourceKeys: + """Phase 1 adds 11 more anthropic_* WIF keys (the anthropic_identity_source discriminator + plus the internal_issuer/keycloak identity-source fields) that need the same dual + registration as the original six tested above.""" + + NEW_KEYS = { + "anthropic_identity_source": "keycloak", + "anthropic_issuer_url": "https://issuer.example", + "anthropic_issuer_subject": "svc-account", + "anthropic_issuer_audience": "https://api.anthropic.com", + "anthropic_issuer_ttl_seconds": "300", + "anthropic_issuer_signing_key_ref": "oidc/env/ISSUER_KEY", + "anthropic_keycloak_token_url": "https://kc.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_auth_method": "client_secret_basic", + "anthropic_keycloak_client_secret_ref": "oidc/env/KC_SECRET", + "anthropic_keycloak_scope": "anthropic-wif", + } + + def test_new_keys_are_exactly_the_non_legacy_registered_set(self): + """Fails the moment a key is added to ANTHROPIC_WIF_KWARGS_KEYS without a matching entry + here (or vice versa), catching drift between what wif.py dispatches on and what this + test (and the funnel/provider-body tests below) actually exercises.""" + from litellm.litellm_core_utils.get_litellm_params import ANTHROPIC_WIF_KWARGS_KEYS + + assert set(self.NEW_KEYS) == ANTHROPIC_WIF_KWARGS_KEYS - set(TestAnthropicWifKeys.SIX_KEYS) + + def test_keys_survive_into_litellm_params(self): + params = get_litellm_params(**self.NEW_KEYS) + for key, value in self.NEW_KEYS.items(): + assert params[key] == value + + def test_keys_are_forwarded_from_completion_kwargs(self): + from litellm.litellm_core_utils.get_litellm_params import FORWARDED_KWARGS_KEYS + + assert set(self.NEW_KEYS) <= FORWARDED_KWARGS_KEYS + + def test_keys_stay_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + for key in self.NEW_KEYS: + assert key in all_litellm_params + + def test_keys_absent_when_not_configured(self): + params = get_litellm_params() + for key in self.NEW_KEYS: + assert key not in params 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 b568b0d91d0..01907901f75 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -2993,6 +2993,20 @@ class TestWifServerOwnedParamsAreUnconditional: "anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_service_account_id", + # Phase 1 identity-source selection and its two variants' fields: each one + # selects a server-side secret or a destination (a signing key, a client + # secret, a token endpoint), so every one joins the same unconditional ban. + "anthropic_identity_source", + "anthropic_issuer_url", + "anthropic_issuer_subject", + "anthropic_issuer_audience", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_signing_key_ref", + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope", ], ) def test_rejected_even_with_proxy_wide_opt_in(self, param: str): @@ -3072,3 +3086,62 @@ class TestWifDisabledOnClientRedirectedBase: assert resolve_anthropic_wif_params({}) is not None assert resolve_anthropic_wif_params({DISABLE_WORKLOAD_IDENTITY_PARAM: True}) is None + + def test_base_override_clears_internal_issuer_fields(self): + """Same failure mode the legacy-path test above guards against, for the internal_issuer + identity source: a signing_key_ref resolved for a client-chosen api_base would mint an + assertion, and then a bearer token, for 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_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": "oidc/env/ISSUER_SIGNING_KEY_PEM", + } + + 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_identity_source" not in redirected + assert "anthropic_issuer_signing_key_ref" not in redirected + assert resolve_anthropic_wif_params(redirected) is None + + def test_base_override_clears_keycloak_fields(self): + """Same as the internal_issuer case above, for the keycloak identity source: a + client_secret_ref resolved for a client-chosen api_base must not follow it there.""" + 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_source": "keycloak", + "anthropic_keycloak_token_url": "https://keycloak.internal.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + + 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_identity_source" not in redirected + assert "anthropic_keycloak_client_secret_ref" not in redirected + assert resolve_anthropic_wif_params(redirected) is None diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_wif.py b/tests/test_litellm/llms/anthropic/test_anthropic_wif.py index 65f38ef87db..56420956ae5 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_wif.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_wif.py @@ -5,7 +5,10 @@ from pathlib import Path from typing import Final import httpx +import jwt import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec import litellm from litellm.llms.anthropic.wif import ( @@ -15,6 +18,12 @@ from litellm.llms.anthropic.wif import ( get_anthropic_wif_token, resolve_anthropic_wif_params, ) +from litellm.llms.base_llm.auth.identity_source import ( + InternalIssuerSource, + KeycloakSource, + identity_source_ref, +) +from litellm.llms.base_llm.auth.jwt_signing import build_jwks, rfc7638_thumbprint from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine from litellm.llms.base_llm.auth.types import ( AssertionSourceError, @@ -32,6 +41,7 @@ WIF_ENV_VARS: Final = ( "ANTHROPIC_WORKSPACE_ID", "ANTHROPIC_IDENTITY_TOKEN_FILE", "ANTHROPIC_IDENTITY_TOKEN", + "ANTHROPIC_IDENTITY_SOURCE", "ANTHROPIC_SCOPE", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL", @@ -534,6 +544,31 @@ class TestErrorMappingExhaustive: assert exc_info.value.llm_provider == "anthropic" assert exc_info.value.model == "claude-sonnet-4-5" + def test_assertion_source_error_detail_is_rendered_when_present(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + _raise_anthropic_wif_error( + AssertionSourceError(kind="unreadable", source_ref="oidc/keycloak/abc123", detail="invalid_client"), + model="claude-sonnet-4-5", + workspace_id_set=True, + ) + + assert "invalid_client" in exc_info.value.message + + def test_assertion_source_error_without_detail_is_unchanged(self): + """Regression floor: the token_file/env path never populates detail, so its message must stay + byte-identical to before the field existed.""" + with pytest.raises(litellm.AuthenticationError) as exc_info: + _raise_anthropic_wif_error( + AssertionSourceError(kind="unreadable", source_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN"), + model="claude-sonnet-4-5", + workspace_id_set=True, + ) + + assert exc_info.value.message == ( + "litellm.AuthenticationError: Anthropic workload identity federation failed. Could not obtain " + "the OIDC identity token (unreadable) from oidc/env/ANTHROPIC_IDENTITY_TOKEN." + ) + def test_endpoint_error_raised_through_facade(self, monkeypatch: pytest.MonkeyPatch): monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]) @@ -635,3 +670,324 @@ class TestFileRereadOnRefresh: assert second == "sk-ant-oat01-second" assert len(poster.requests) == 2 assert poster.requests[1].json_body()["assertion"] == "second-assertion" + + +_ISSUER_PRIVATE_VALUE: Final = 55566677788899900011122233344455566677788899900011122233344455 +ISSUER_SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +KEYCLOAK_TOKEN_URL: Final = "https://keycloak.internal.example/realms/litellm/protocol/openid-connect/token" + + +def _issuer_signing_key() -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(_ISSUER_PRIVATE_VALUE, ec.SECP256R1()) + + +def _issuer_signing_key_pem() -> str: + return ( + _issuer_signing_key() + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + +def _get_secret_str_returning(pem: str, ref: str) -> Callable[..., str | None]: + def fake_get_secret_str(secret_name: str, default_value: str | None = None) -> str | None: + return pem if secret_name == ref else default_value + + return fake_get_secret_str + + +class TestIdentitySourceDiscriminatorAbsentIsByteIdenticalToLegacy: + """anthropic_identity_source unset must resolve exactly like today: no new dispatch code + runs, and no assertion_source closure is attached, so the engine falls back to its own + reader precisely as it always has.""" + + def test_file_config_carries_no_assertion_source(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": str(token_file), + } + ) + + assert params == AnthropicWifParams( + federation_rule_id="fdrl_1", + organization_id="org-1", + assertion_ref=f"oidc/file/{token_file}", + ) + assert params.assertion_source is None + + def test_env_config_carries_no_assertion_source(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt") + + params = resolve_anthropic_wif_params(None) + + assert params is not None + assert params.assertion_ref == "oidc/env/ANTHROPIC_IDENTITY_TOKEN" + assert params.assertion_source is None + + +class TestInternalIssuerIdentitySourceDispatch: + """A config.yaml-shaped litellm_params block for the internal_issuer identity source.""" + + LITELLM_PARAMS: Final = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_ttl_seconds": 300, + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + } + + def test_assertion_ref_matches_the_identity_source_hash(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + expected_config = InternalIssuerSource( + issuer_url="https://issuer.internal.example", + subject="workload-a", + ttl_seconds=300, + signing_key_ref=ISSUER_SIGNING_KEY_REF, + ) + assert params.assertion_ref == identity_source_ref(expected_config) + assert params.assertion_ref.startswith("oidc/internal_issuer/") + + def test_ref_is_stable_and_rolls_on_field_change(self): + first = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + second = resolve_anthropic_wif_params(dict(self.LITELLM_PARAMS)) + changed = resolve_anthropic_wif_params({**self.LITELLM_PARAMS, "anthropic_issuer_subject": "workload-b"}) + + assert first is not None and second is not None and changed is not None + assert first.assertion_ref == second.assertion_ref + assert first.assertion_ref != changed.assertion_ref + + def test_assertion_source_mints_a_verifiable_jwt(self, monkeypatch: pytest.MonkeyPatch): + pem = _issuer_signing_key_pem() + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + _get_secret_str_returning(pem, ISSUER_SIGNING_KEY_REF), + ) + + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + assert params is not None + assert params.assertion_source is not None + + assertion = params.assertion_source() + + assert assertion is not None + public_key = _issuer_signing_key().public_key() + expected_kid = build_jwks(public_key)["keys"][0]["kid"] + assert jwt.get_unverified_header(assertion)["kid"] == expected_kid + assert expected_kid == rfc7638_thumbprint(public_key) + claims = jwt.decode(assertion, public_key, algorithms=["ES256"], options={"verify_aud": False}) + assert claims["sub"] == "workload-a" + assert claims["iss"] == "https://issuer.internal.example" + + def test_full_exchange_sends_the_minted_assertion(self, monkeypatch: pytest.MonkeyPatch): + pem = _issuer_signing_key_pem() + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + _get_secret_str_returning(pem, ISSUER_SIGNING_KEY_REF), + ) + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + token = get_anthropic_wif_token(self.LITELLM_PARAMS, "https://api.anthropic.com", "claude-sonnet-4-5", engine) + + assert token == "sk-ant-oat01-minted" + sent_assertion = poster.requests[0].json_body()["assertion"] + jwt.decode( + sent_assertion, _issuer_signing_key().public_key(), algorithms=["ES256"], options={"verify_aud": False} + ) + + +class TestKeycloakIdentitySourceDispatch: + """A config.yaml-shaped litellm_params block for the keycloak identity source. The minted + closure's own network behavior is covered by test_client_credentials.py's DI-poster tests; + this only proves wif.py threads the fields into the right config and hash.""" + + LITELLM_PARAMS: Final = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": KEYCLOAK_TOKEN_URL, + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + + def test_assertion_ref_matches_the_identity_source_hash(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + expected_config = KeycloakSource( + token_url=KEYCLOAK_TOKEN_URL, + client_id="litellm", + client_secret_ref="oidc/env/KEYCLOAK_CLIENT_SECRET", + ) + assert params.assertion_ref == identity_source_ref(expected_config) + assert params.assertion_ref.startswith("oidc/keycloak/") + + def test_assertion_source_is_a_fresh_closure(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + assert params.assertion_source is not None + assert callable(params.assertion_source) + + def test_auth_method_change_rolls_the_ref(self): + default_method = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + post_method = resolve_anthropic_wif_params( + {**self.LITELLM_PARAMS, "anthropic_keycloak_auth_method": "client_secret_post"} + ) + + assert default_method is not None and post_method is not None + assert default_method.assertion_ref != post_method.assertion_ref + + def test_client_secret_ref_pointer_name_change_rolls_the_ref_without_resolving_it(self): + """The hash covers the pointer NAME, never a resolved secret (decision 7) -- true even + though nothing in this test ever calls get_secret_str.""" + first = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + second = resolve_anthropic_wif_params( + {**self.LITELLM_PARAMS, "anthropic_keycloak_client_secret_ref": "oidc/env/OTHER_SECRET_NAME"} + ) + + assert first is not None and second is not None + assert first.assertion_ref != second.assertion_ref + + +class TestIdentitySourceValidationFailsClosed: + """Unknown discriminator, a missing required variant field, and a field belonging to the + other variant are all hard config errors at resolution time -- never a silent fallback to + token_file (decision 5).""" + + def test_unknown_discriminator_raises(self): + with pytest.raises(litellm.AuthenticationError, match="anthropic_identity_source"): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "bogus", + } + ) + + def test_internal_issuer_missing_required_fields_raises(self): + with pytest.raises(litellm.AuthenticationError): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + } + ) + + def test_keycloak_missing_required_fields_raises(self): + with pytest.raises(litellm.AuthenticationError): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_client_id": "litellm", + } + ) + + def test_mixed_variant_fields_raise(self): + with pytest.raises(litellm.AuthenticationError, match="belongs to a different identity source"): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + "anthropic_keycloak_client_id": "leaked-from-other-variant", + } + ) + + def test_secret_pasted_into_wrong_field_never_appears_in_the_error(self): + secret_value = "super-secret-client-value-xyz" + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + "anthropic_issuer_ttl_seconds": secret_value, + } + ) + + assert secret_value not in exc_info.value.message + + +class TestConfigYamlShapedIdentitySources: + """One litellm_params dict per identity source, shaped exactly like the + model_list[].litellm_params block a proxy config.yaml carries -- proving an operator can + configure each of Phase 1's supported sources.""" + + def test_legacy_token_file_source(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_token_file": str(token_file), + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref == f"oidc/file/{token_file}" + assert params.assertion_source is None + + def test_internal_issuer_source(self): + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://litellm.internal.example", + "anthropic_issuer_subject": "litellm-proxy", + "anthropic_issuer_ttl_seconds": 300, + "anthropic_issuer_signing_key_ref": "os.environ/ISSUER_SIGNING_KEY_PEM", + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref.startswith("oidc/internal_issuer/") + assert params.assertion_source is not None + + def test_keycloak_source(self): + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": KEYCLOAK_TOKEN_URL, + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_auth_method": "client_secret_post", + "anthropic_keycloak_client_secret_ref": "os.environ/KEYCLOAK_CLIENT_SECRET", + "anthropic_keycloak_scope": "anthropic-wif", + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref.startswith("oidc/keycloak/") + assert params.assertion_source is not None diff --git a/tests/test_litellm/llms/base_llm/auth/test_client_credentials.py b/tests/test_litellm/llms/base_llm/auth/test_client_credentials.py new file mode 100644 index 00000000000..c9bffb9c010 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/auth/test_client_credentials.py @@ -0,0 +1,311 @@ +import base64 +import logging +from collections.abc import Mapping +from typing import Final +from urllib.parse import parse_qsl + +import httpx +import pytest + +from litellm.llms.base_llm.auth.client_credentials import ( + fetch_keycloak_assertion, + keycloak_assertion_source, +) +from litellm.llms.base_llm.auth.identity_source import KeycloakSource, identity_source_ref + +TOKEN_URL: Final = "https://keycloak.example/realms/litellm/protocol/openid-connect/token" +CLIENT_ID: Final = "litellm" +CLIENT_SECRET_REF: Final = "oidc/env/KEYCLOAK_CLIENT_SECRET" +CLIENT_SECRET: Final = "s3cr3t-client-value" + + +class RecordedRequest: + def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None: + self.url = url + self.content = content + self.headers = dict(headers) + self.timeout = timeout + + def form_body(self) -> dict[str, str]: + return dict(parse_qsl(self.content.decode())) + + +class ScriptedPoster: + """Returns one scripted response per call; records every request it receives.""" + + def __init__(self, responses: list[httpx.Response]) -> None: + self.requests: list[RecordedRequest] = [] + self._responses = list(responses) + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.requests.append(RecordedRequest(url, content, headers, timeout)) + return self._responses.pop(0) if len(self._responses) > 1 else self._responses[0] + + +class RaisingPoster: + def __init__(self, error: Exception) -> None: + self.calls = 0 + self._error = error + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + raise self._error + + +def make_config( + auth_method: str = "client_secret_basic", + scope: str | None = None, + token_url: str = TOKEN_URL, + client_secret_ref: str = CLIENT_SECRET_REF, +) -> KeycloakSource: + return KeycloakSource( + token_url=token_url, + client_id=CLIENT_ID, + client_secret_ref=client_secret_ref, + auth_method=auth_method, # pyright: ignore[reportArgumentType] # test-only string widened for parametrization + scope=scope, + ) + + +def secret_reader_returning(secret: str | None): + def reader(ref: str) -> str | None: + assert ref == CLIENT_SECRET_REF + return secret + + return reader + + +DEFAULT_SECRET_READER: Final = secret_reader_returning(CLIENT_SECRET) + + +def token_response(access_token: str = "keycloak-minted-token") -> httpx.Response: + return httpx.Response(200, json={"access_token": access_token, "token_type": "Bearer", "expires_in": 300}) + + +class TestClientSecretBasic: + def test_sends_basic_auth_header_and_no_secret_in_body(self): + poster = ScriptedPoster([token_response("minted-1")]) + + token = fetch_keycloak_assertion( + make_config(auth_method="client_secret_basic"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert token == "minted-1" + request = poster.requests[0] + assert request.url == TOKEN_URL + expected_auth = "Basic " + base64.b64encode(f"{CLIENT_ID}:{CLIENT_SECRET}".encode()).decode("ascii") + assert request.headers["authorization"] == expected_auth + assert request.headers["content-type"] == "application/x-www-form-urlencoded" + body = request.form_body() + assert body["grant_type"] == "client_credentials" + assert "client_secret" not in body + assert "client_id" not in body + + def test_scope_included_only_when_set(self): + poster = ScriptedPoster([token_response()]) + fetch_keycloak_assertion( + make_config(scope="openid profile"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert poster.requests[0].form_body()["scope"] == "openid profile" + + poster_no_scope = ScriptedPoster([token_response()]) + fetch_keycloak_assertion(make_config(scope=None), poster=poster_no_scope, secret_reader=DEFAULT_SECRET_READER) + + assert "scope" not in poster_no_scope.requests[0].form_body() + + +class TestClientSecretPost: + def test_sends_client_id_and_secret_in_body_with_no_basic_header(self): + poster = ScriptedPoster([token_response("minted-2")]) + + token = fetch_keycloak_assertion( + make_config(auth_method="client_secret_post"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert token == "minted-2" + request = poster.requests[0] + assert "authorization" not in request.headers + body = request.form_body() + assert body["grant_type"] == "client_credentials" + assert body["client_id"] == CLIENT_ID + assert body["client_secret"] == CLIENT_SECRET + + +class TestOnePostPerExchange: + def test_exactly_one_post_per_call_no_cache(self): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + + first = fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + second = fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert first == "first" + assert second == "second" + assert len(poster.requests) == 2 + + +class TestInvalidClient: + def test_400_invalid_client_surfaces_redacted_detail(self): + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": "unauthorized client"})] + ) + + with pytest.raises(ValueError, match="invalid_client") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert "unauthorized client" in str(exc_info.value) + assert "400" in str(exc_info.value) + assert CLIENT_SECRET not in str(exc_info.value) + + def test_echoed_client_secret_is_never_reflected_into_the_error(self): + """A misbehaving Keycloak that echoes the submitted client_secret back in its error body + must never leak it into the exception the caller sees.""" + long_secret: Final = "reflectable-secret-0123456789" + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": f"got {long_secret} in body"})] + ) + + with pytest.raises(ValueError, match="keycloak") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(long_secret)) + + assert long_secret not in str(exc_info.value) + assert "redacted" in str(exc_info.value) + + def test_echoed_short_client_secret_is_never_reflected_into_the_error(self): + """Real Keycloak client secrets are often shorter than a JWT: the reflection probe must + not silently stop protecting a secret just because it is under the probe's usual length.""" + short_secret: Final = "hand-set-14ch" + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": f"got {short_secret} in body"})] + ) + + with pytest.raises(ValueError, match="keycloak") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(short_secret)) + + assert short_secret not in str(exc_info.value) + assert "redacted" in str(exc_info.value) + + +class TestUnreachable: + def test_transport_failure_raises_diagnosable_value_error(self): + poster = RaisingPoster(httpx.ConnectError("connection refused")) + + with pytest.raises(ValueError, match="ConnectError") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert poster.calls == 1 + assert CLIENT_SECRET not in str(exc_info.value) + + +class TestNon2xx: + def test_500_raises_value_error_with_status_code(self): + poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]) + + with pytest.raises(ValueError, match="500"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + +class TestResponseValidation: + def test_missing_access_token_is_a_value_error(self): + poster = ScriptedPoster([httpx.Response(200, json={"token_type": "Bearer"})]) + + with pytest.raises(ValueError, match="schema validation"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + def test_empty_access_token_is_a_value_error(self): + poster = ScriptedPoster([httpx.Response(200, json={"access_token": " "})]) + + with pytest.raises(ValueError, match="empty access_token"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + +class TestInsecureTokenUrl: + def test_http_url_is_rejected_before_any_post(self): + poster = ScriptedPoster([token_response()]) + + with pytest.raises(ValueError, match="https"): + fetch_keycloak_assertion( + make_config(token_url="http://keycloak.example/token"), + poster=poster, + secret_reader=DEFAULT_SECRET_READER, + ) + + assert poster.requests == [] + + +class TestMissingClientSecret: + def test_unresolvable_secret_ref_raises_value_error_naming_the_ref_not_a_secret(self): + poster = ScriptedPoster([token_response()]) + + with pytest.raises(ValueError, match=CLIENT_SECRET_REF): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(None)) + + assert poster.requests == [] + + +class TestKeycloakAssertionSource: + def test_returns_a_callable_that_fetches_fresh_each_call(self): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + source = keycloak_assertion_source(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert source() == "first" + assert source() == "second" + assert len(poster.requests) == 2 + + def test_propagates_the_underlying_fetch_failure(self): + poster = ScriptedPoster([httpx.Response(400, json={"error": "invalid_client"})]) + source = keycloak_assertion_source(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + with pytest.raises(ValueError, match="invalid_client"): + source() + + +class TestClientSecretNeverLeaks: + """Regression coverage for the load-bearing property: a Keycloak client_secret must never + surface in the assertion_ref, in any error message, or in a log record, however it fails.""" + + def test_never_in_the_assertion_ref(self): + config = make_config(client_secret_ref=CLIENT_SECRET_REF) + + ref = identity_source_ref(config) + + assert CLIENT_SECRET not in ref + assert CLIENT_SECRET_REF not in ref + + def test_never_in_any_raised_error_message_across_every_failure_mode(self): + config = make_config() + failures = [ + lambda: fetch_keycloak_assertion( + config, + poster=ScriptedPoster([httpx.Response(400, json={"error": "invalid_client"})]), + secret_reader=DEFAULT_SECRET_READER, + ), + lambda: fetch_keycloak_assertion( + config, poster=RaisingPoster(httpx.ConnectError("boom")), secret_reader=DEFAULT_SECRET_READER + ), + lambda: fetch_keycloak_assertion( + config, + poster=ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]), + secret_reader=DEFAULT_SECRET_READER, + ), + lambda: fetch_keycloak_assertion( + config, poster=ScriptedPoster([token_response()]), secret_reader=secret_reader_returning(None) + ), + ] + for fail in failures: + with pytest.raises(ValueError, match="keycloak") as exc_info: + fail() + assert CLIENT_SECRET not in str(exc_info.value) + + def test_never_in_a_log_record(self, caplog: pytest.LogCaptureFixture): + with caplog.at_level(logging.DEBUG): + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": CLIENT_SECRET})] + ) + with pytest.raises(ValueError, match="keycloak"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + fetch_keycloak_assertion( + make_config(), poster=ScriptedPoster([token_response()]), secret_reader=DEFAULT_SECRET_READER + ) + + assert CLIENT_SECRET not in caplog.text diff --git a/tests/test_litellm/llms/base_llm/auth/test_identity_source.py b/tests/test_litellm/llms/base_llm/auth/test_identity_source.py new file mode 100644 index 00000000000..bfa8847c69a --- /dev/null +++ b/tests/test_litellm/llms/base_llm/auth/test_identity_source.py @@ -0,0 +1,239 @@ +from types import MappingProxyType +from typing import Final, Literal + +import pytest +from pydantic import ValidationError + +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_config_adapter, + identity_source_ref, +) + +SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +OTHER_SIGNING_KEY_REF: Final = "oidc/env/OTHER_SIGNING_KEY_PEM" +CLIENT_SECRET_REF: Final = "oidc/env/KEYCLOAK_CLIENT_SECRET" +ISSUER_URL: Final = "https://issuer.internal.example" +SUBJECT: Final = "workload-a" +TOKEN_URL: Final = "https://keycloak.example/realms/litellm/protocol/openid-connect/token" +CLIENT_ID: Final = "litellm" + + +def make_issuer( + issuer_url: str = ISSUER_URL, + subject: str = SUBJECT, + signing_key_ref: str = SIGNING_KEY_REF, + ttl_seconds: int = 300, +) -> InternalIssuerSource: + return InternalIssuerSource( + issuer_url=issuer_url, subject=subject, signing_key_ref=signing_key_ref, ttl_seconds=ttl_seconds + ) + + +def make_keycloak( + token_url: str = TOKEN_URL, + client_id: str = CLIENT_ID, + client_secret_ref: str = CLIENT_SECRET_REF, + auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic", + scope: str | None = None, +) -> KeycloakSource: + return KeycloakSource( + token_url=token_url, + client_id=client_id, + client_secret_ref=client_secret_ref, + auth_method=auth_method, + scope=scope, + ) + + +class TestIdentitySourceRefHashing: + def test_identical_config_hashes_idempotently(self): + assert identity_source_ref(make_issuer()) == identity_source_ref(make_issuer()) + + def test_ref_is_prefixed_by_kind(self): + assert identity_source_ref(make_issuer()).startswith("oidc/internal_issuer/") + assert identity_source_ref(make_keycloak()).startswith("oidc/keycloak/") + + def test_pointer_name_change_changes_ref(self): + """Two configs differing only in which secret a pointer names must never collide, since a + stale ref would let the token exchange's outer cache key alias two different credentials.""" + first: Final = identity_source_ref(make_issuer(signing_key_ref=SIGNING_KEY_REF)) + second: Final = identity_source_ref(make_issuer(signing_key_ref=OTHER_SIGNING_KEY_REF)) + + assert first != second + + def test_non_pointer_field_change_changes_ref(self): + first: Final = identity_source_ref(make_keycloak(scope="openid")) + second: Final = identity_source_ref(make_keycloak(scope="openid profile")) + + assert first != second + + def test_ref_never_contains_the_pointer_field_values(self): + """The ref is a fixed-width hash, not a serialization of the config, so no field value - + pointer name or otherwise - can leak into the secret-free string echoed into errors.""" + ref: Final = identity_source_ref(make_issuer()) + + assert SIGNING_KEY_REF not in ref + assert "issuer.internal.example" not in ref + + def test_different_kinds_with_disjoint_fields_never_collide(self): + assert identity_source_ref(make_issuer()) != identity_source_ref(make_keycloak()) + + +class TestInternalIssuerSourceValidation: + def test_defaults(self): + source: Final = make_issuer() + + assert source.kind == AnthropicIdentitySourceKind.internal_issuer + assert source.ttl_seconds == 300 + assert source.audience is None + + def test_ttl_seconds_over_one_hour_is_rejected(self): + with pytest.raises(ValidationError): + make_issuer(ttl_seconds=3601) + + def test_ttl_seconds_at_one_hour_is_accepted(self): + assert make_issuer(ttl_seconds=3600).ttl_seconds == 3600 + + def test_non_positive_ttl_seconds_is_rejected(self): + with pytest.raises(ValidationError): + make_issuer(ttl_seconds=0) + + def test_missing_signing_key_ref_is_rejected(self): + missing_field: Final = MappingProxyType({"issuer_url": ISSUER_URL, "subject": SUBJECT}) + + with pytest.raises(ValidationError): + InternalIssuerSource.model_validate(missing_field) + + def test_keycloak_only_field_is_rejected_as_extra(self): + mixed_variant: Final = MappingProxyType( + { + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + + with pytest.raises(ValidationError): + InternalIssuerSource.model_validate(mixed_variant) + + def test_is_frozen(self): + source: Final = make_issuer() + + with pytest.raises(ValidationError): + source.subject = "workload-b" + + def test_secret_pasted_into_wrong_typed_field_is_not_echoed_in_the_error(self): + """hide_input_in_errors keeps a value the operator pasted into a mistyped field out of the + validation error, so a client_secret headed for the wrong field isn't logged in the raise.""" + leaked_secret: Final = "shh-do-not-log-me" + wrong_type: Final = MappingProxyType( + { + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "ttl_seconds": leaked_secret, + } + ) + + with pytest.raises(ValidationError) as exc_info: + InternalIssuerSource.model_validate(wrong_type) + + assert leaked_secret not in str(exc_info.value) + + +class TestKeycloakSourceValidation: + def test_defaults(self): + source: Final = make_keycloak() + + assert source.kind == AnthropicIdentitySourceKind.keycloak + assert source.auth_method == "client_secret_basic" + assert source.scope is None + + def test_client_secret_post_is_accepted(self): + assert make_keycloak(auth_method="client_secret_post").auth_method == "client_secret_post" + + def test_private_key_jwt_is_not_a_supported_auth_method_yet(self): + unshipped_auth_method: Final = MappingProxyType( + { + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + "auth_method": "private_key_jwt", + } + ) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(unshipped_auth_method) + + def test_audience_field_was_dropped(self): + dropped_field: Final = MappingProxyType( + { + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + "audience": "https://anthropic.example", + } + ) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(dropped_field) + + def test_missing_client_secret_ref_is_rejected(self): + missing_field: Final = MappingProxyType({"token_url": TOKEN_URL, "client_id": CLIENT_ID}) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(missing_field) + + +class TestDiscriminatedUnionParsing: + def test_parses_internal_issuer_variant(self): + parsed: Final = identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "internal_issuer", + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + } + ) + ) + + assert isinstance(parsed, InternalIssuerSource) + + def test_parses_keycloak_variant(self): + parsed: Final = identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "keycloak", + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + ) + + assert isinstance(parsed, KeycloakSource) + + def test_unknown_kind_is_a_hard_error(self): + with pytest.raises(ValidationError): + identity_source_config_adapter.validate_python(MappingProxyType({"kind": "token_file"})) + + def test_mixed_variant_fields_are_a_hard_error(self): + """A keycloak field on an internal_issuer-tagged payload must fail closed rather than be + silently dropped or silently accepted as if it selected the other variant.""" + with pytest.raises(ValidationError): + identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "internal_issuer", + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + ) diff --git a/tests/test_litellm/llms/base_llm/auth/test_internal_issuer.py b/tests/test_litellm/llms/base_llm/auth/test_internal_issuer.py new file mode 100644 index 00000000000..d965a356426 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/auth/test_internal_issuer.py @@ -0,0 +1,188 @@ +import json +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec + +from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource +from litellm.llms.base_llm.auth.internal_issuer import ( + internal_issuer_assertion_source, + internal_issuer_jwks_document, + mint_internal_issuer_assertion, +) +from litellm.llms.base_llm.auth.jwt_signing import build_jwks, rfc7638_thumbprint + +SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +ISSUER_URL: Final = "https://issuer.internal.example" +SUBJECT: Final = "workload-a" + + +_PRIVATE_VALUE: Final = 90123456789012345678901234567890123456789012345678901234567890 + + +def signing_key() -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(_PRIVATE_VALUE, ec.SECP256R1()) + + +def pem_of(key: ec.EllipticCurvePrivateKey) -> str: + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +def make_config( + issuer_url: str = ISSUER_URL, + subject: str = SUBJECT, + audience: str | None = None, + ttl_seconds: int = 300, + signing_key_ref: str = SIGNING_KEY_REF, +) -> InternalIssuerSource: + return InternalIssuerSource( + issuer_url=issuer_url, + subject=subject, + audience=audience, + ttl_seconds=ttl_seconds, + signing_key_ref=signing_key_ref, + ) + + +def key_reader_returning(pem: str | None): + def reader(ref: str) -> str | None: + assert ref == SIGNING_KEY_REF + return pem + + return reader + + +class FakeClock: + def __init__(self, value: float) -> None: + self._value: Final = value + + def __call__(self) -> float: + return self._value + + +def decode_ignoring_wall_clock(token: str, public_key: ec.EllipticCurvePublicKey) -> dict: + """Tests mint with a fixed past ``FakeClock`` and no expected audience, so PyJWT's + real-wall-clock ``exp``/``aud`` checks (irrelevant to what these tests verify) are disabled.""" + return jwt.decode(token, public_key, algorithms=["ES256"], options={"verify_exp": False, "verify_aud": False}) + + +class TestMintInternalIssuerAssertion: + def test_required_claims_and_asymmetric_alg(self): + key: Final = signing_key() + config: Final = make_config(ttl_seconds=300) + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + header: Final = jwt.get_unverified_header(token) + claims: Final = decode_ignoring_wall_clock(token, key.public_key()) + + assert header["alg"] == "ES256" + assert claims["sub"] == SUBJECT + assert claims["iss"] == ISSUER_URL + assert claims["iat"] == 1_700_000_000 + assert claims["exp"] == 1_700_000_300 + + def test_kid_matches_the_published_jwks(self): + key: Final = signing_key() + config: Final = make_config() + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + + header_kid: Final = jwt.get_unverified_header(token)["kid"] + published_kid: Final = build_jwks(key.public_key())["keys"][0]["kid"] + assert header_kid == published_kid == rfc7638_thumbprint(key.public_key()) + + def test_ttl_bounds_exp_minus_iat(self): + key: Final = signing_key() + config: Final = make_config(ttl_seconds=120) + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + claims: Final = decode_ignoring_wall_clock(token, key.public_key()) + + assert claims["exp"] - claims["iat"] == 120 + + def test_audience_included_only_when_set(self): + key: Final = signing_key() + without_audience: Final = mint_internal_issuer_assertion( + make_config(audience=None), key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + with_audience: Final = mint_internal_issuer_assertion( + make_config(audience="urn:anthropic:federation"), + key_reader=key_reader_returning(pem_of(key)), + clock=FakeClock(1_700_000_000.0), + ) + + claims_without: Final = decode_ignoring_wall_clock(without_audience, key.public_key()) + claims_with: Final = decode_ignoring_wall_clock(with_audience, key.public_key()) + assert "aud" not in claims_without + assert claims_with["aud"] == "urn:anthropic:federation" + + def test_jti_is_present_and_fresh_on_every_mint(self): + key: Final = signing_key() + config: Final = make_config() + reader: Final = key_reader_returning(pem_of(key)) + + first: Final = decode_ignoring_wall_clock( + mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)), + key.public_key(), + ) + second: Final = decode_ignoring_wall_clock( + mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)), + key.public_key(), + ) + + assert first["jti"] and second["jti"] + assert first["jti"] != second["jti"] + + def test_missing_signing_key_raises_value_error_naming_the_ref_not_a_secret(self): + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + mint_internal_issuer_assertion(make_config(), key_reader=key_reader_returning(None)) + + def test_malformed_signing_key_raises_value_error(self): + with pytest.raises(ValueError, match="not a valid unencrypted PEM"): + mint_internal_issuer_assertion(make_config(), key_reader=key_reader_returning("not-a-pem")) + + +class TestInternalIssuerAssertionSource: + def test_returns_a_callable_that_mints_fresh_each_call(self): + key: Final = signing_key() + source: Final = internal_issuer_assertion_source(make_config(), key_reader=key_reader_returning(pem_of(key))) + + first: Final = jwt.decode(source(), key.public_key(), algorithms=["ES256"]) + second: Final = jwt.decode(source(), key.public_key(), algorithms=["ES256"]) + + assert first["jti"] != second["jti"] + + def test_propagates_the_underlying_mint_failure(self): + source: Final = internal_issuer_assertion_source(make_config(), key_reader=key_reader_returning(None)) + + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + source() + + +class TestInternalIssuerJwksDocument: + def test_matches_the_key_used_to_mint(self): + key: Final = signing_key() + config: Final = make_config() + reader: Final = key_reader_returning(pem_of(key)) + + document: Final = json.loads(internal_issuer_jwks_document(config, key_reader=reader)) + token: Final = mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)) + + assert document["keys"][0]["kid"] == jwt.get_unverified_header(token)["kid"] + assert decode_ignoring_wall_clock(token, key.public_key()) + + def test_missing_signing_key_raises_value_error(self): + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + internal_issuer_jwks_document(make_config(), key_reader=key_reader_returning(None)) diff --git a/tests/test_litellm/llms/base_llm/auth/test_jwt_signing.py b/tests/test_litellm/llms/base_llm/auth/test_jwt_signing.py new file mode 100644 index 00000000000..77c303ef4e2 --- /dev/null +++ b/tests/test_litellm/llms/base_llm/auth/test_jwt_signing.py @@ -0,0 +1,174 @@ +import base64 +import hashlib +import json +import time +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from litellm.llms.base_llm.auth.jwt_signing import ( + build_jwk, + build_jwks, + jwks_document_json, + load_es256_private_key, + rfc7638_thumbprint, + sign_es256_jwt, +) + +_FIXED_PRIVATE_VALUE: Final = 55090612345678901234567890123456789012345678901234567890123456 +_OTHER_PRIVATE_VALUE: Final = 1 + + +def fixed_private_key(value: int = _FIXED_PRIVATE_VALUE) -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(value, ec.SECP256R1()) + + +def pem_of(key: ec.EllipticCurvePrivateKey) -> str: + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +def independent_thumbprint(public_key: ec.EllipticCurvePublicKey) -> str: + """Recomputes RFC 7638 by hand, deliberately not sharing a single line of code with + ``jwt_signing.rfc7638_thumbprint`` -- a mutation that broke the real implementation must not + also break this reference, or the two would trivially agree by sharing the bug.""" + numbers: Final = public_key.public_numbers() + x: Final = base64.urlsafe_b64encode(numbers.x.to_bytes(32, "big")).rstrip(b"=").decode() + y: Final = base64.urlsafe_b64encode(numbers.y.to_bytes(32, "big")).rstrip(b"=").decode() + canonical: Final = f'{{"crv":"P-256","kty":"EC","x":"{x}","y":"{y}"}}' + return base64.urlsafe_b64encode(hashlib.sha256(canonical.encode()).digest()).rstrip(b"=").decode() + + +class TestLoadEs256PrivateKey: + def test_valid_ec_p256_pem_loads(self): + key: Final = load_es256_private_key(pem_of(fixed_private_key())) + + assert isinstance(key, ec.EllipticCurvePrivateKey) + assert isinstance(key.curve, ec.SECP256R1) + + def test_garbage_pem_is_rejected(self): + with pytest.raises(ValueError, match="not a valid unencrypted PEM"): + load_es256_private_key("not a pem") + + def test_rsa_key_is_rejected(self): + rsa_pem: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(rsa_pem) + + def test_non_p256_curve_is_rejected(self): + secp384_pem: Final = ( + ec.generate_private_key(ec.SECP384R1()) + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(secp384_pem) + + def test_error_never_echoes_key_material(self): + pem: Final = pem_of(fixed_private_key()) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(pem_of(ec.generate_private_key(ec.SECP384R1()))) + with pytest.raises(ValueError, match="not a valid unencrypted PEM") as exc_info: + load_es256_private_key("garbage-not-a-pem") + + assert pem not in str(exc_info.value) + + +class TestRfc7638Thumbprint: + def test_matches_independent_recomputation(self): + public_key: Final = fixed_private_key().public_key() + + assert rfc7638_thumbprint(public_key) == independent_thumbprint(public_key) + + def test_different_keys_have_different_thumbprints(self): + first: Final = fixed_private_key(_FIXED_PRIVATE_VALUE).public_key() + second: Final = fixed_private_key(_OTHER_PRIVATE_VALUE).public_key() + + assert rfc7638_thumbprint(first) != rfc7638_thumbprint(second) + + def test_thumbprint_is_deterministic(self): + public_key: Final = fixed_private_key().public_key() + + assert rfc7638_thumbprint(public_key) == rfc7638_thumbprint(public_key) + + +class TestBuildJwks: + def test_jwks_contains_one_key_matching_the_thumbprint(self): + public_key: Final = fixed_private_key().public_key() + + jwks: Final = build_jwks(public_key) + + assert len(jwks["keys"]) == 1 + assert jwks["keys"][0]["kid"] == rfc7638_thumbprint(public_key) + assert jwks["keys"][0]["kty"] == "EC" + assert jwks["keys"][0]["crv"] == "P-256" + assert jwks["keys"][0]["alg"] == "ES256" + + def test_build_jwk_stamps_the_given_kid_verbatim(self): + jwk: Final = build_jwk(fixed_private_key().public_key(), kid="caller-supplied-kid") + + assert jwk["kid"] == "caller-supplied-kid" + + def test_jwks_document_json_round_trips_through_build_jwks(self): + key: Final = fixed_private_key() + + document: Final = json.loads(jwks_document_json(pem_of(key))) + jwks: Final = build_jwks(key.public_key()) + + assert document == {"keys": [dict(jwk) for jwk in jwks["keys"]]} + + +class TestSignEs256Jwt: + def test_minted_token_verifies_against_the_matching_public_key(self): + key: Final = fixed_private_key() + now: Final = int(time.time()) + claims: Final = {"sub": "workload-a", "iss": "https://issuer.example", "iat": now, "exp": now + 300} + + token: Final = sign_es256_jwt(pem_of(key), claims) + decoded: Final = jwt.decode(token, key.public_key(), algorithms=["ES256"]) + + assert decoded == claims + + def test_header_alg_is_es256(self): + token: Final = sign_es256_jwt(pem_of(fixed_private_key()), {"sub": "x"}) + + assert jwt.get_unverified_header(token)["alg"] == "ES256" + + def test_header_kid_matches_the_published_jwks(self): + key: Final = fixed_private_key() + + token: Final = sign_es256_jwt(pem_of(key), {"sub": "x"}) + + header_kid: Final = jwt.get_unverified_header(token)["kid"] + published_kid: Final = build_jwks(key.public_key())["keys"][0]["kid"] + assert header_kid == published_kid == rfc7638_thumbprint(key.public_key()) + + def test_wrong_key_fails_verification(self): + signing_key: Final = fixed_private_key(_FIXED_PRIVATE_VALUE) + other_key: Final = fixed_private_key(_OTHER_PRIVATE_VALUE) + + token: Final = sign_es256_jwt(pem_of(signing_key), {"sub": "x"}) + + with pytest.raises(jwt.exceptions.InvalidSignatureError): + jwt.decode(token, other_key.public_key(), algorithms=["ES256"]) diff --git a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py index c4fe3c221c7..d808d384356 100644 --- a/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py +++ b/tests/test_litellm/llms/base_llm/auth/test_token_exchange.py @@ -9,10 +9,11 @@ from typing import Final from urllib.parse import parse_qsl import httpx -from pydantic import SecretStr import pytest +from pydantic import SecretStr from litellm.llms.base_llm.auth.token_exchange import ( + _REDACTION_CAP, ADVISORY_REFRESH_BACKOFF_SECONDS, FALLBACK_TOKEN_TTL_SECONDS, MAX_ASSERTION_BYTES, @@ -629,6 +630,108 @@ class TestAssertionGuards: assert result.kind == expected_kind assert len(poster.requests) == 0 + def test_value_error_message_is_captured_as_detail(self): + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise ValueError("Keycloak token endpoint returned invalid_client") + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail == "Keycloak token endpoint returned invalid_client" + + @pytest.mark.parametrize( + "raised", + [OidcPathNotAllowedError("path outside allowed credential directories"), OSError("permission denied")], + ) + def test_non_value_error_never_populates_detail(self, raised: Exception): + """Only the ValueError branch carries operator-diagnosable text; every other reader failure + stays detail=None, matching today's file/env behavior byte-for-byte.""" + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise raised + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail is None + + def test_value_error_detail_is_capped(self): + poster = ScriptedPoster([token_response()]) + overlong_message = "x" * (_REDACTION_CAP + 100) + + def reader(ref: str) -> str | None: + raise ValueError(overlong_message) + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail == overlong_message[:_REDACTION_CAP] + assert len(result.detail) == _REDACTION_CAP + + +class TestAssertionSourceOverridesEngineReader: + """``TokenExchangeSpec.assertion_source`` is the dispatch mechanism a per-config identity + source (internal_issuer, keycloak) plugs into the shared engine with -- it must win over the + engine-level reader, and failures must still be reported against ``assertion_ref``.""" + + def test_assertion_source_is_used_instead_of_the_reader(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + spec = make_spec(assertion_source=lambda: "from-assertion-source") + + result = mint(engine, spec) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert poster.requests[0].json_body()["assertion"] == "from-assertion-source" + + def test_reader_is_never_called_when_assertion_source_is_set(self): + poster = ScriptedPoster([token_response()]) + calls: list[str] = [] + + def reader(ref: str) -> str | None: + calls.append(ref) + return "from-engine-reader" + + engine = make_engine(poster, reader=reader) + spec = make_spec(assertion_source=lambda: "from-assertion-source") + + mint(engine, spec) + + assert calls == [] + + def test_assertion_source_failure_is_reported_against_assertion_ref(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + + def raising_source() -> str | None: + raise ValueError("keycloak token endpoint returned invalid_client") + + spec = make_spec(assertion_source=raising_source, assertion_ref="oidc/keycloak/abc123") + + result = engine.get_token(spec) + + assert isinstance(result, AssertionSourceError) + assert result.source_ref == "oidc/keycloak/abc123" + assert result.detail == "keycloak token endpoint returned invalid_client" + assert len(poster.requests) == 0 + + def test_assertion_source_is_re_invoked_on_401_retry(self): + """The 401-retry re-read (``_reread_assertion``) must also prefer ``assertion_source``, + not silently fall back to the engine reader for the reflected-assertion check.""" + values = iter(["assertion-v1", "assertion-v2"]) + poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"}), token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + spec = make_spec(assertion_source=lambda: next(values)) + + result = mint(engine, spec) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert poster.requests[0].json_body()["assertion"] == "assertion-v1" + assert poster.requests[1].json_body()["assertion"] == "assertion-v2" + class TestOidcFilePathAllowlistRaisesTypedError: """The engine classifies assertion-source failures by exception type (see diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 9301176f3ed..f0dbf952e8f 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -29,6 +29,20 @@ from litellm.proxy.auth.auth_utils import ( ) +def test_every_anthropic_wif_kwarg_key_is_request_banned(): + """anthropic_wif_litellm_params (types/utils.py) is derived from ANTHROPIC_WIF_KWARGS_KEYS + (get_litellm_params.py) precisely so a new WIF field can never be added to the kwargs funnel + without automatically joining the request-body ban list; this guards that invariant itself, + independent of today's field count, so it fails if the derivation is ever reverted to a + hand-typed list that drifts.""" + 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) + + class TestCustomAuthCommonChecksWarning: """custom_auth_common_checks_warning only warns when custom auth is configured and the common-checks opt-in is off, since that is the only state where diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index d00fbf589e3..9af106122c3 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4542,6 +4542,40 @@ def test_get_deployment_credentials_with_provider_preserves_aws_auth_params(): assert credentials.get(key) == value, key +def test_get_deployment_credentials_with_provider_preserves_anthropic_wif_params(): + """ + Test that get_deployment_credentials_with_provider preserves a litellm_params-configured + Anthropic workload identity federation setup (both the legacy token_file fields and the + Phase 1 internal_issuer/keycloak identity-source fields) so files/batches/passthrough + deployments using WIF do not silently fall back to a missing credential. + """ + wif_params = { + "anthropic_federation_rule_id": "fdrl_deployment", + "anthropic_organization_id": "org-deployment", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": "https://keycloak.internal.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + router = litellm.Router( + model_list=[ + { + "model_name": "anthropic-wif-model", + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + **wif_params, + }, + } + ], + ) + + credentials = router.get_deployment_credentials_with_provider(model_id="anthropic-wif-model") + + assert credentials is not None + for key, value in wif_params.items(): + assert credentials.get(key) == value, key + + def _team_wildcard_model(api_key: str, model_id: str = "team-wildcard-id") -> dict: return { "model_name": f"model_name_team-1_{model_id}", diff --git a/tests/test_litellm/types/test_router.py b/tests/test_litellm/types/test_router.py index accd3b32a0d..aedc6c078d6 100644 --- a/tests/test_litellm/types/test_router.py +++ b/tests/test_litellm/types/test_router.py @@ -2,11 +2,16 @@ import pytest from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, + CredentialLiteLLMParams, Deployment, LiteLLM_Params, ModelInfo, ) -from litellm.types.utils import CustomPricingLiteLLMParams, MirroredPricingParams +from litellm.types.utils import ( + CustomPricingLiteLLMParams, + MirroredPricingParams, + anthropic_wif_litellm_params, +) def test_model_info_declares_mirrored_pricing_fields(): @@ -89,3 +94,21 @@ def test_pricing_strings_are_coerced_to_float(): def test_invalid_pricing_is_rejected(): with pytest.raises(ValueError, match='validation error for ModelInfo'): ModelInfo(id="x", input_cost_per_token="free") + + +def test_credential_litellm_params_declares_every_anthropic_wif_field(): + """Without these, get_deployment_credentials_with_provider round-trips litellm_params + through a strict Pydantic dump and silently drops every WIF field before files/batches/ + passthrough callers see it -- the same #30235-shaped gap azure_ad_token closed above.""" + for field in anthropic_wif_litellm_params: + assert field in CredentialLiteLLMParams.model_fields, field + + +def test_anthropic_wif_fields_round_trip_through_model_dump(): + values = {field: f"value-for-{field}" for field in anthropic_wif_litellm_params} + values["anthropic_issuer_ttl_seconds"] = 300 + + dumped = CredentialLiteLLMParams(**values).model_dump(exclude_none=True) + + for field, value in values.items(): + assert dumped[field] == value, field