mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(anthropic): pluggable identity sources for workload identity federation
The workload assertion could only come from a mounted file or an environment variable, which assumes a platform that already projects one. Two more sources sit behind an explicit anthropic_identity_source discriminator, and its absence keeps today's resolver exactly as it was: - internal_issuer signs a short-lived ES256 assertion with an operator-supplied key, resolved through the usual secret reference so it can live in a secret manager. Its JWKS is exported for the operator to register with Anthropic, without which the source cannot be used at all. - keycloak fetches the assertion from a client_credentials grant, with the client secret likewise held behind a reference rather than in the configuration The engine gains an optional assertion source on the spec, so a source that needs more than a string can supply one without the reference ever carrying a secret: it stays a hash of the non-secret fields and remains what the cache keys on and what errors name. Failures carry a redacted detail, so a Keycloak hop is diagnosable rather than collapsing into one opaque message The server-owned field list is now derived from one definition, so every field added here is rejected from request bodies and cleared on a client base override without a second edit. Credential params carry the federation fields too, which is what the files, batches, and passthrough surfaces read Verified against the live token endpoint: an ES256 assertion from internal_issuer, with its exported JWKS registered on the federation issuer, mints a token. Note that a freshly registered inline JWKS takes up to about a minute to become usable
This commit is contained in:
parent
79e4a6936d
commit
66a37443c3
23 changed files with 2250 additions and 30 deletions
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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/<kind>/<hash>`` 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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
187
litellm/llms/base_llm/auth/client_credentials.py
Normal file
187
litellm/llms/base_llm/auth/client_credentials.py
Normal file
|
|
@ -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)
|
||||
56
litellm/llms/base_llm/auth/identity_source.py
Normal file
56
litellm/llms/base_llm/auth/identity_source.py
Normal file
|
|
@ -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/<kind>/<hash>`` 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/<kind>/<hash>``: 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}"
|
||||
84
litellm/llms/base_llm/auth/internal_issuer.py
Normal file
84
litellm/llms/base_llm/auth/internal_issuer.py
Normal file
|
|
@ -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))
|
||||
103
litellm/llms/base_llm/auth/jwt_signing.py
Normal file
103
litellm/llms/base_llm/auth/jwt_signing.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__"})
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
311
tests/test_litellm/llms/base_llm/auth/test_client_credentials.py
Normal file
311
tests/test_litellm/llms/base_llm/auth/test_client_credentials.py
Normal file
|
|
@ -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
|
||||
239
tests/test_litellm/llms/base_llm/auth/test_identity_source.py
Normal file
239
tests/test_litellm/llms/base_llm/auth/test_identity_source.py
Normal file
|
|
@ -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,
|
||||
}
|
||||
)
|
||||
)
|
||||
188
tests/test_litellm/llms/base_llm/auth/test_internal_issuer.py
Normal file
188
tests/test_litellm/llms/base_llm/auth/test_internal_issuer.py
Normal file
|
|
@ -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))
|
||||
174
tests/test_litellm/llms/base_llm/auth/test_jwt_signing.py
Normal file
174
tests/test_litellm/llms/base_llm/auth/test_jwt_signing.py
Normal file
|
|
@ -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"])
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue