diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index abfb9dc56fe..43f32938553 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -211,6 +211,7 @@ from litellm.types.mcp_server.mcp_server_manager import ( MCPServer, parse_pinned_tools, ) +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm from litellm.types.utils import CallTypes if TYPE_CHECKING: @@ -419,7 +420,7 @@ class MCPServerConfig(TypedDict, total=False): id_jag_resource: str client_private_key: str client_private_key_id: str - client_assertion_signing_alg: str + client_assertion_signing_alg: ApprovedJwtAlgorithm timeout: float max_concurrent_requests: int @@ -1567,6 +1568,18 @@ def _warn_on_shared_identifier_prefixes(servers: Iterable[MCPServer]) -> None: ) +def _stored_client_assertion_signing_alg(value: object, server_name: str) -> ApprovedJwtAlgorithm: + if value in APPROVED_JWT_ALGORITHMS: + return cast(ApprovedJwtAlgorithm, value) # cast-ok: value was checked against APPROVED_JWT_ALGORITHMS + if value is not None: + verbose_logger.warning( + "MCP server %s: client_assertion_signing_alg %r is not an approved algorithm, using RS256", + server_name, + value, + ) + return "RS256" + + def _warn_legacy_delegate_auth_if_applicable(server: MCPServer, *, source: str) -> None: """Direct legacy delegated OAuth configurations to the admitted replacement.""" if server.auth_type != MCPAuth.oauth2: @@ -2707,7 +2720,10 @@ class MCPServerManager: id_jag_resource=server_config.get("id_jag_resource", None), client_private_key=server_config.get("client_private_key", None), client_private_key_id=server_config.get("client_private_key_id", None), - client_assertion_signing_alg=server_config.get("client_assertion_signing_alg", "RS256"), + client_assertion_signing_alg=_stored_client_assertion_signing_alg( + server_config.get("client_assertion_signing_alg"), + str(server_config.get("alias") or server_config.get("server_name") or server_id), + ), token_exchange_profile=server_config.get("token_exchange_profile", "rfc8693"), allow_sampling=bool(server_config.get("allow_sampling", False)), allow_elicitation=bool(server_config.get("allow_elicitation", False)), @@ -3290,10 +3306,10 @@ class MCPServerManager: credentials_are_encrypted, ), client_private_key_id=(credentials_dict.get("client_private_key_id") if credentials_dict else None), - client_assertion_signing_alg=( - credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None - ) - or "RS256", + client_assertion_signing_alg=_stored_client_assertion_signing_alg( + credentials_dict.get("client_assertion_signing_alg") if credentials_dict else None, + name_for_prefix, + ), token_exchange_profile=mcp_server.token_exchange_profile or (credentials_dict.get("token_exchange_profile") if credentials_dict else None) or "rfc8693", diff --git a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py index c31560a9f63..0c8c3ceeb39 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py +++ b/litellm/proxy/_experimental/mcp_server/oauth_identity_binding.py @@ -23,20 +23,12 @@ from typing_extensions import assert_never from litellm._logging import verbose_logger from litellm.caching.in_memory_cache import InMemoryCache from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.proxy.auth.jwt_algorithms import jwks_keys_for from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.mcp_server.mcp_server_manager import MCPOAuthIdentityBinding, MCPServer +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS -_ALLOWED_ID_TOKEN_ALGORITHMS: Final = ( - "RS256", - "RS384", - "RS512", - "ES256", - "ES384", - "ES512", - "PS256", - "PS384", - "PS512", -) +_ALLOWED_ID_TOKEN_ALGORITHMS: Final = APPROVED_JWT_ALGORITHMS _JWKS_CACHE_TTL_SECONDS: Final = 3600 _jwks_cache: Final = InMemoryCache(default_ttl=_JWKS_CACHE_TTL_SECONDS) @@ -126,7 +118,7 @@ async def _discover_jwks_url(issuer: str) -> str: def _select_signing_key(id_token: str, keys: Sequence[Mapping[str, object]]) -> "jwt.PyJWK | _BindingRejection": header: Final = jwt.get_unverified_header(id_token) kid: Final = header.get("kid") - for key in keys: + for key in jwks_keys_for(keys, _ALLOWED_ID_TOKEN_ALGORITHMS): if kid is None or key.get("kid") == kid: return jwt.PyJWK(dict(key)) return _BindingRejection( diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index 33c3a854058..8a7f2653971 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -45,6 +45,7 @@ from litellm.types.mcp import ( DEFAULT_SUBJECT_TOKEN_TYPE, normalize_upstream_header_name, ) +from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm class AuthResolution(str, Enum): @@ -326,7 +327,7 @@ class PrivateKeyJwtAuth(BaseModel): source: Literal["private_key_jwt"] = "private_key_jwt" private_key: SecretStr key_id: str | None = None - signing_alg: str = "RS256" + signing_alg: ApprovedJwtAlgorithm = "RS256" class ClientSecretAuth(BaseModel): diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0028838b673..0eb1df1c860 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -17,6 +17,7 @@ from pydantic import ( PositiveInt, PrivateAttr, TypeAdapter, + ValidationError, field_validator, model_validator, ) @@ -52,6 +53,10 @@ from litellm.types.mcp import ( ) from litellm.types.mcp_server.mcp_server_manager import MCPInfo from litellm.types.proxy.agent_identity import ManagedAgentContext +from litellm.types.proxy.auth.jwt_algorithms import ( + APPROVED_JWT_ALGORITHMS, + ApprovedJwtAlgorithm, +) from litellm.types.proxy.carried_budget_state import ( OrgBudgetSnapshot, TeamBudgetSnapshot, @@ -1593,6 +1598,19 @@ def _reject_unsupported_per_server_oauth_discovery(values: object, require_auth_ raise _per_server_oauth_discovery_error() +_APPROVED_JWT_ALGORITHM_ADAPTER: Final = TypeAdapter(ApprovedJwtAlgorithm) + + +def _validated_client_assertion_signing_alg(credentials: MCPCredentials | None) -> MCPCredentials | None: + if credentials is None or credentials.get("client_assertion_signing_alg") is None: + return credentials + try: + _APPROVED_JWT_ALGORITHM_ADAPTER.validate_python(credentials["client_assertion_signing_alg"]) + except ValidationError as exc: + raise ValueError(f"client_assertion_signing_alg must be one of {', '.join(APPROVED_JWT_ALGORITHMS)}") from exc + return credentials + + def _validate_mcp_transport_fields(values: object) -> None: if not isinstance(values, dict): return @@ -1680,6 +1698,11 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): description="Server-managed: set by the endpoint; caller values are overridden.", ) + @field_validator("credentials") + @classmethod + def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None: + return _validated_client_assertion_signing_alg(credentials) + @model_validator(mode="after") def validate_protocol_transport(self) -> "NewMCPServerRequest": validate_mcp_protocol_transport( @@ -1773,6 +1796,11 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): timeout: float | None = None max_concurrent_requests: int | None = None + @field_validator("credentials") + @classmethod + def check_client_assertion_signing_alg(cls, credentials: MCPCredentials | None) -> MCPCredentials | None: + return _validated_client_assertion_signing_alg(credentials) + @model_validator(mode="after") def validate_protocol_transport(self) -> "UpdateMCPServerRequest": if not {"transport", "mcp_info"}.issubset(self.model_fields_set): diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index cfa9d10b74c..8dac077fe3d 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -16,6 +16,7 @@ import re import time from collections.abc import Awaitable, Callable, Collection, Mapping, Sequence from dataclasses import dataclass +from functools import lru_cache from typing import Any, Final, Literal, NoReturn, Protocol, TypeVar, cast import httpx @@ -58,6 +59,7 @@ from litellm.proxy.agent_endpoints.identity import has_legacy_identity from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore, resolve_managed_agent from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failure from litellm.proxy.auth.auth_checks import can_team_access_model +from litellm.proxy.auth.jwt_algorithms import allowed_jwt_algorithms, jwks_keys_for from litellm.proxy.auth.model_access_denied import ( ModelAccessDeniedHTTPException, model_access_denied_client_message, @@ -65,6 +67,7 @@ from litellm.proxy.auth.model_access_denied import ( from litellm.proxy.auth.resolvers.grants import GrantResolver, UserLookup, canonical_user_id from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.team_grants import team_grants, team_model_aliases +from litellm.proxy.common_utils.fips import is_fips_mode from litellm.proxy.common_utils.user_api_key_cache import ( AUTH_OBJECTS_TARGET, UserApiKeyCache, @@ -75,6 +78,7 @@ from litellm.repositories.user_repository import UserRepository from litellm.types.agents import AgentResponse from litellm.types.proxy.agent_identity import AgentIdentityFailure from litellm.types.proxy.auth.auth_checks import UserNotFoundError +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS from .auth_checks import ( TeamNotFoundError, @@ -111,6 +115,14 @@ UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_" _CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str) +@lru_cache(maxsize=1) +def _log_eddsa_deprecation() -> None: + verbose_proxy_logger.warning( + "JWT Auth: accepted a token signed with EdDSA, which is deprecated and not FIPS 140-3 approved; " + "it is rejected when LITELLM_FIPS_MODE=true. Move the IdP signing key to RS256, PS256 or ES256" + ) + + class _JWTAuthSettings(Protocol): """The JWT auth settings block this handler reads back through ``getattr``, when one is configured.""" @@ -222,16 +234,8 @@ class JWTHandler: # "Warning: Make sure not to mix symmetric and asymmetric algorithms that interpret # the key in different ways (e.g. HS* and RS*)." SUPPORTED_JWT_ALGORITHMS = [ - "RS256", - "RS384", - "RS512", - "PS256", - "PS384", - "PS512", - "ES256", - "ES384", - "ES512", - "EdDSA", + *APPROVED_JWT_ALGORITHMS, + *LEGACY_JWT_ALGORITHMS, ] LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer" LITELLM_USER_ID_CLAIM = "_litellm_user_id" @@ -252,13 +256,24 @@ class JWTHandler: def __init__( self, + fips_mode: Callable[[], bool] = is_fips_mode, ) -> None: + self._fips_mode: Final = fips_mode self.http_handler = HTTPHandler() self.leeway = 0 # Per-cache-key locks so a TTL lapse triggers one refresh instead of one per in-flight request. self._refresh_locks: dict[str, asyncio.Lock] = {} # mutable-ok: lock registry, keyed by JWKS url self.agent_lookup: AgentLookup = _NoRegisteredAgents() + def allowed_algorithms(self) -> tuple[str, ...]: + return allowed_jwt_algorithms(self._fips_mode()) + + def _warn_deprecated_signing_algorithm(self, token: str) -> None: + if self._fips_mode(): + return + if jwt.get_unverified_header(token).get("alg") == "EdDSA": + _log_eddsa_deprecation() + def bind_agent_lookup(self, agent_lookup: AgentLookup) -> None: self.agent_lookup = agent_lookup @@ -958,7 +973,13 @@ class JWTHandler: log_context=f"kid={kid}", ) - public_key: Final = self.parse_keys(keys=keys, kid=kid) + allowed: Final = self.allowed_algorithms() + usable_keys: Final[JWKKeyValue] = ( + list(jwks_keys_for(keys, allowed)) + if isinstance(keys, list) + else next(iter(jwks_keys_for((keys,), allowed)), {}) + ) + public_key: Final = self.parse_keys(keys=usable_keys, kid=kid) if public_key is not None: return cast(dict, public_key) @@ -1239,32 +1260,27 @@ class JWTHandler: ) ) - if isinstance(public_key, dict): - public_key_obj: Final = PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key - return jwt.decode( - token, - public_key_obj, - algorithms=self.SUPPORTED_JWT_ALGORITHMS, - options=decode_options, - audience=audience, - issuer=issuer, - leeway=self.leeway, + key_obj: Final = ( + PyJWK.from_dict(self._get_jwk_from_public_key(public_key=public_key)).key + if isinstance(public_key, dict) + else x509.load_pem_x509_certificate(public_key.encode(), default_backend()) + .public_key() + .public_bytes( + serialization.Encoding.PEM, + serialization.PublicFormat.SubjectPublicKeyInfo, ) - - cert: Final = x509.load_pem_x509_certificate(public_key.encode(), default_backend()) - key: Final = cert.public_key().public_bytes( - serialization.Encoding.PEM, - serialization.PublicFormat.SubjectPublicKeyInfo, ) - return jwt.decode( + payload: Final = jwt.decode( token, - key, - algorithms=self.SUPPORTED_JWT_ALGORITHMS, + key_obj, + algorithms=self.allowed_algorithms(), audience=audience, issuer=issuer, options=decode_options, leeway=self.leeway, ) + self._warn_deprecated_signing_algorithm(token) + return payload async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict: try: diff --git a/litellm/proxy/auth/jwt_algorithms.py b/litellm/proxy/auth/jwt_algorithms.py new file mode 100644 index 00000000000..0e1839ec200 --- /dev/null +++ b/litellm/proxy/auth/jwt_algorithms.py @@ -0,0 +1,35 @@ +from collections.abc import Collection, Mapping, Sequence +from types import MappingProxyType +from typing import Final + +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, LEGACY_JWT_ALGORITHMS + +_KEY_TYPE_ALGORITHMS: Final = MappingProxyType( + { + "RSA": frozenset(APPROVED_JWT_ALGORITHMS[:6]), + "EC": frozenset(APPROVED_JWT_ALGORITHMS[6:]), + "OKP": frozenset(LEGACY_JWT_ALGORITHMS), + } +) + + +def allowed_jwt_algorithms(fips_mode: bool) -> tuple[str, ...]: + return APPROVED_JWT_ALGORITHMS if fips_mode else (*APPROVED_JWT_ALGORITHMS, *LEGACY_JWT_ALGORITHMS) + + +def jwks_keys_for( + keys: Sequence[Mapping[str, object]], algorithms: Collection[str] +) -> tuple[Mapping[str, object], ...]: + """Keep keys whose declared alg is allowed; keys without alg are kept only when their kty can sign with an + allowed algorithm.""" + return tuple(key for key in keys if _key_allowed(key, frozenset(algorithms))) + + +def _key_allowed(key: Mapping[str, object], algorithms: frozenset[str]) -> bool: + alg: Final = key.get("alg") + if isinstance(alg, str): + return alg in algorithms + key_type: Final = key.get("kty") + if not isinstance(key_type, str): + return False + return bool(_KEY_TYPE_ALGORITHMS.get(key_type, frozenset()) & algorithms) diff --git a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py index 221c4b3752b..75ce60ebacf 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/mcp_jwt_signer/mcp_jwt_signer.py @@ -90,8 +90,10 @@ from litellm.integrations.custom_guardrail import ( log_guardrail_information, ) from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.jwt_algorithms import jwks_keys_for from litellm.types.guardrail_base_init import GuardrailBaseInitKwargs from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS from litellm.types.utils import CallTypesLiteral if TYPE_CHECKING: @@ -452,12 +454,16 @@ class MCPJWTSigner(CustomGuardrail): unverified_header: Final = jwt.get_unverified_header(raw_token) kid: Final = unverified_header.get("kid") + approved_keys: Final = jwks_keys_for(jwks_keys, APPROVED_JWT_ALGORITHMS) + if not approved_keys: + raise jwt.exceptions.PyJWKSetError(f"No JWKS key at {jwks_uri!r} uses an approved signing algorithm") + # Build a JWKS object and pick the matching key. # PyJWT's PyJWKSet handles key-type parsing and kid matching correctly. from jwt import PyJWKSet try: - jwks_set: Final = PyJWKSet.from_dict({"keys": jwks_keys}) + jwks_set: Final = PyJWKSet.from_dict({"keys": list(approved_keys)}) except Exception as exc: raise jwt.exceptions.PyJWKSetError(f"Failed to parse JWKS from {jwks_uri!r}: {exc}") from exc diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 343fb731b06..cee6ec09f32 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -15,6 +15,7 @@ from litellm.types.mcp import ( normalize_upstream_header_name, validate_mcp_protocol_transport, ) +from litellm.types.proxy.auth.jwt_algorithms import ApprovedJwtAlgorithm # MCPInfo now allows arbitrary additional fields for custom metadata @@ -178,7 +179,7 @@ class MCPServer(BaseModel): id_jag_resource: str | None = None client_private_key: str | None = None client_private_key_id: str | None = None - client_assertion_signing_alg: str = "RS256" + client_assertion_signing_alg: ApprovedJwtAlgorithm = "RS256" # Wire dialect: "rfc8693" (standard token-exchange grant) or "entra_obo" (Microsoft Entra # On-Behalf-Of, the RFC 7523 jwt-bearer grant + requested_token_use extension) token_exchange_profile: str = "rfc8693" diff --git a/litellm/types/proxy/auth/jwt_algorithms.py b/litellm/types/proxy/auth/jwt_algorithms.py new file mode 100644 index 00000000000..73c383dd9a1 --- /dev/null +++ b/litellm/types/proxy/auth/jwt_algorithms.py @@ -0,0 +1,17 @@ +from typing import Final, Literal + +ApprovedJwtAlgorithm = Literal["RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512"] + +APPROVED_JWT_ALGORITHMS: Final[tuple[ApprovedJwtAlgorithm, ...]] = ( + "RS256", + "RS384", + "RS512", + "PS256", + "PS384", + "PS512", + "ES256", + "ES384", + "ES512", +) + +LEGACY_JWT_ALGORITHMS: Final = ("EdDSA",) diff --git a/tests/integration/authorization/test_jwt_algorithm_allowlist.py b/tests/integration/authorization/test_jwt_algorithm_allowlist.py new file mode 100644 index 00000000000..44a5df0505e --- /dev/null +++ b/tests/integration/authorization/test_jwt_algorithm_allowlist.py @@ -0,0 +1,721 @@ +import asyncio +import base64 +import json +import os +import socket +import time +import uuid +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final + +import httpx +import jwt +import psutil +import yaml +from anthropic import Anthropic +from cryptography.hazmat.primitives.asymmetric import ec, ed25519, rsa +from jwt.algorithms import ECAlgorithm, OKPAlgorithm, RSAAlgorithm +from openai import AsyncOpenAI, OpenAI + +from tests.integration._support.client import Gateway, eventually +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_proxy_process +from tests.integration._support.wire import Reply, Request, wire_server + +KEY_ID: Final = "integration-jwt-allowlist-key" +APPROVED_WARNING_PARTS: Final = ("EdDSA", "deprecated", "LITELLM_FIPS_MODE") + + +def _jwks_reply(public_jwk: str) -> Reply: + return _jwks_reply_keys([{**json.loads(public_jwk), "kid": KEY_ID}]) + + +def _jwks_reply_keys(keys: list[dict[str, object]]) -> Reply: + return Reply(body=json.dumps({"keys": keys}).encode()) + + +def _eddsa_keypair() -> tuple[ed25519.Ed25519PrivateKey, dict[str, object]]: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + return private_key, {**json.loads(OKPAlgorithm.to_jwk(private_key.public_key())), "kid": "ed"} + + +def _rsa_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return private_key, {**json.loads(RSAAlgorithm.to_jwk(private_key.public_key())), "kid": "rsa"} + + +def _ec_keypair() -> tuple[ec.EllipticCurvePrivateKey, dict[str, object]]: + private_key: Final = ec.generate_private_key(ec.SECP256R1()) + return private_key, {**json.loads(ECAlgorithm.to_jwk(private_key.public_key())), "kid": "ec"} + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def _observed_bodies() -> tuple[str, ...]: + response: Final = httpx.get(f"{os.environ['INTEGRATION_UPSTREAM_URL']}/__observations", timeout=15) + assert response.status_code == 200, response.text + return tuple(json.dumps(entry["body"]) for entry in response.json()["requests"]) + + +def _jwt_auth_config(tmp_path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config["general_settings"], + "enable_jwt_auth": True, + "litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True}, + } + path: Final = tmp_path / "jwt_allowlist.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _token(key: object, algorithm: str, subject: str, kid: str = KEY_ID) -> str: + return jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + key, # pyright: ignore[reportArgumentType] # jwt.encode takes Any key material + algorithm=algorithm, + headers={"kid": kid}, + ) + + +def test_eddsa_signed_token_is_accepted_and_logged_as_deprecated_outside_fips_mode( + gateway: Gateway, tmp_path: Path +) -> None: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + public_jwk: Final = OKPAlgorithm.to_jwk(private_key.public_key()) + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply(public_jwk) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path) + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: "EdDSA" in text and "deprecated" in text and "LITELLM_FIPS_MODE" in text, + seconds=30, + ) + assert "EdDSA" in deprecation and "deprecated" in deprecation and "LITELLM_FIPS_MODE" in deprecation + + +def test_rs256_signed_token_is_accepted_without_deprecation_log(gateway: Gateway, tmp_path: Path) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = RSAAlgorithm.to_jwk(private_key.public_key()) + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply(public_jwk) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path) + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + assert "EdDSA" not in owned.log.read_text(), owned.log.read_text() + + +def _deprecation_lines(log_text: str) -> list[str]: + return [line for line in log_text.splitlines() if all(part in line for part in APPROVED_WARNING_PARTS)] + + +def _alg_lied_token(private_key: ed25519.Ed25519PrivateKey, kid: str, subject: str) -> str: + def segment(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode() + + signing_input: Final = ( + f"{segment(json.dumps({'alg': 'RS256', 'typ': 'JWT', 'kid': kid}).encode())}." + f"{segment(json.dumps({'sub': subject, 'iat': int(time.time()), 'exp': int(time.time()) + 300}).encode())}" + ) + signature: Final = segment(private_key.sign(signing_input.encode())) + return f"{signing_input}.{signature}" + + +def test_eddsa_non_stream_request_reaches_upstream_and_logs_deprecation(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + marker: Final = f"jwt-allowlist-a1-{uuid.uuid4().hex}" + _observed_bodies() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}]}, + key=token, + ) + assert response.status_code == 200, response.text + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + assert _deprecation_lines(deprecation) != [] + observed: Final = _observed_bodies() + assert any(marker in body for body in observed), observed + + +def test_eddsa_streamed_request_through_openai_sdk_reaches_upstream_and_logs_deprecation( + gateway: Gateway, tmp_path: Path +) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + marker: Final = f"jwt-allowlist-a3-{uuid.uuid4().hex}" + _observed_bodies() + sdk: Final = OpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30) + chunks: Final = [ + chunk + for chunk in sdk.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True + ) + ] + assert chunks, "stream produced no chunks" + deprecation: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + assert _deprecation_lines(deprecation) != [] + observed: Final = _observed_bodies() + assert any(marker in body for body in observed), observed + + +def test_rs256_messages_request_through_anthropic_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _rsa_keypair() + completion: Final = json.dumps( + { + "id": "msg_allowlist_a4", + "type": "message", + "role": "assistant", + "model": "claude-3-5-haiku-20241022", + "content": [{"type": "text", "text": "a4-reply"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 2}, + } + ).encode() + + def jwks(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + def provider(request: Request) -> Reply: + return Reply(body=completion) + + with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject, kid="rsa") + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": keys_server.url}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model(model="anthropic/claude-3-5-haiku-latest", api_base=provider_server.url) + scenario.cleanups.callback(scenario.delete_user, subject) + sdk: Final = Anthropic( + base_url=str(owned.gateway.client.base_url).rstrip("/"), + auth_token=token, + timeout=30, + max_retries=0, + ) + message: Final = sdk.messages.create( + model=model, max_tokens=8, messages=[{"role": "user", "content": "a4"}] + ) + text: Final = message.content[0].text + assert "a4-reply" in text, message + received: Final = provider_server.drain() + assert received, "upstream never received the request" + + +def test_rs256_responses_request_through_openai_async_sdk_reaches_upstream(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _rsa_keypair() + response_object: Final = json.dumps( + { + "id": "resp_allowlist_a5", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_a5", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "a5-reply", "annotations": []}], + } + ], + "usage": {"input_tokens": 5, "output_tokens": 2, "total_tokens": 7}, + } + ).encode() + + def jwks(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + def provider(request: Request) -> Reply: + return Reply(body=response_object) + + with wire_server(jwks) as keys_server, wire_server(provider) as provider_server, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "RS256", subject, kid="rsa") + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": keys_server.url}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model(api_base=provider_server.url) + scenario.cleanups.callback(scenario.delete_user, subject) + + async def call() -> str: + sdk: Final = AsyncOpenAI(base_url=f"{owned.gateway.client.base_url}/v1", api_key=token, timeout=30) + response: Final = await sdk.responses.create(model=model, input="a5") + return response.id + + identity: Final = asyncio.run(call()) + assert identity.startswith("resp_"), identity + received: Final = provider_server.drain() + assert received, "upstream never received the request" + + +def test_es256_signed_token_is_accepted(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _ec_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + private_key, + algorithm="ES256", + headers={"kid": "ec"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist es256"}]}, + key=token, + ) + assert response.status_code == 200, response.text + + +def test_hs256_token_against_oct_jwks_key_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + secret: Final = b"integration-hs256-jwt-secret-0123456789abcdef" + oct_key: Final = { + "kty": "oct", + "kid": "sym", + "alg": "HS256", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([oct_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + secret, + algorithm="HS256", + headers={"kid": "sym"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist hs256"}]}, + key=token, + ) + assert response.status_code == 401, response.text + assert "error" in response.text, response.text + + +def test_token_signed_eddsa_with_rs256_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _alg_lied_token(private_key, "ed", subject) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist lied"}]}, + key=token, + ) + assert response.status_code == 401, response.text + + +def test_request_without_authorization_header_is_rejected(gateway: Gateway, tmp_path: Path) -> None: + _, jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.client.post( + "/v1/chat/completions", + json={"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist noauth"}]}, + ) + assert response.status_code == 401, response.text + + +def test_empty_jwks_document_rejects_tokens(gateway: Gateway, tmp_path: Path) -> None: + private_key, _jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + private_key, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist empty"}]}, + key=token, + ) + assert response.status_code == 401, response.text + + +def test_five_eddsa_requests_warn_at_most_once_per_worker(gateway: Gateway, tmp_path: Path) -> None: + private_key, jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = _token(private_key, "EdDSA", subject, kid="ed") + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + for _ in range(5): + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist burst"}]}, + key=token, + ) + assert response.status_code == 200, response.text + log_text: Final = eventually( + lambda: owned.log.read_text(), + lambda text: _deprecation_lines(text) != [], + seconds=30, + ) + warnings: Final = _deprecation_lines(log_text) + assert 1 <= len(warnings) <= 2, warnings + + +def test_master_key_request_still_works_on_jwt_enabled_proxy(gateway: Gateway, tmp_path: Path) -> None: + _, jwks_key = _rsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + response: Final = owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt alg allowlist master"}]}, + ) + assert response.status_code == 200, response.text + + +def _response_ids(status: int, text: str) -> tuple[str, ...]: + if status != 200: + return () + stripped: Final = text.strip() + if stripped.startswith("{"): + body: Final = json.loads(stripped) + return (str(body["id"]),) if "id" in body else () + ids: Final = { + str(chunk["id"]) + for line in stripped.splitlines() + if line.startswith("data: ") and line[len("data: ") :].strip() != "[DONE]" + for chunk in (json.loads(line[len("data: ") :]),) + if "id" in chunk + } + return tuple(ids) + + +def _spend_row_counts(identifiers: set[str]) -> dict[str, int]: + if not identifiers: + return {} + rows: Final = read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])', + ("{" + ",".join(sorted(identifiers)) + "}",), + ) + counts: Final = {identifier: 0 for identifier in identifiers} + for row in rows: + identifier = str(row["request_id"]) + if identifier in counts: + counts[identifier] += 1 + return counts + + +def test_mixed_jwt_burst_survives_jwks_restart_and_writes_spend_rows(gateway: Gateway, tmp_path: Path) -> None: + rsa_private, rsa_jwks_key = _rsa_keypair() + ed_private, ed_jwks_key = _eddsa_keypair() + port: Final = _free_port() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key]) + + with gateway.scenario() as scenario: + rs_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + rsa_private, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + ed_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + ed_private, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + with owned_proxy_process( + gateway, + tmp_path, + {"JWT_PUBLIC_KEY_URL": f"http://127.0.0.1:{port}/jwks"}, + config=_jwt_auth_config(tmp_path), + workers=2, + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"]) + ) + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"]) + ) + + def chat(token: str, index: int, stream: bool = False) -> httpx.Response: + with owned.gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": f"d1-{index}"}], + **({"stream": True} if stream else {}), + }, + headers={"Authorization": f"Bearer {token}"}, + ) as response: + return httpx.Response( + status_code=response.status_code, + headers=dict(response.headers), + content=response.read(), + ) + + with wire_server(respond, port=port): + warmup_rs: Final = chat(rs_token, 0) + warmup_ed: Final = chat(ed_token, 1) + assert warmup_rs.status_code == 200, warmup_rs.text + assert warmup_ed.status_code == 200, warmup_ed.text + + statuses: Final[list[int]] = [] + bodies: Final[list[str]] = [] + with ThreadPoolExecutor(max_workers=15) as pool: + futures = [ + pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index, index % 5 == 0) + for index in range(30) + ] + for future in futures: + result = future.result(timeout=60) + statuses.append(result.status_code) + bodies.append(result.text) + assert all(status in (200, 401) for status in statuses), (statuses, bodies[:3]) + assert any(status == 200 for status in statuses), statuses + + with wire_server(respond, port=port): + fresh_rs: Final = chat(rs_token, 100) + fresh_ed: Final = chat(ed_token, 101) + assert fresh_rs.status_code == 200, fresh_rs.text + assert fresh_ed.status_code == 200, fresh_ed.text + + succeeded: Final = { + identifier for status, text in zip(statuses, bodies) for identifier in _response_ids(status, text) + } + succeeded.update(_response_ids(200, fresh_rs.text) + _response_ids(200, fresh_ed.text)) + assert len(succeeded) == sum(1 for status in statuses if status == 200) + 2, ( + statuses, + bodies[:3], + ) + counts: Final = eventually( + lambda: _spend_row_counts(succeeded), + lambda value: value == {identifier: 1 for identifier in succeeded}, + seconds=70, + ) + assert all(count == 1 for count in counts.values()), counts + + +def test_burst_survives_killed_worker_and_stays_alive(gateway: Gateway, tmp_path: Path) -> None: + rsa_private, rsa_jwks_key = _rsa_keypair() + ed_private, ed_jwks_key = _eddsa_keypair() + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply_keys([rsa_jwks_key, ed_jwks_key]) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + rs_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + rsa_private, + algorithm="RS256", + headers={"kid": "rsa"}, + ) + ed_token: Final = jwt.encode( + {"sub": f"integration-jwt-{uuid.uuid4().hex}", "iat": int(time.time()), "exp": int(time.time()) + 900}, + ed_private, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + with owned_proxy_process( + gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=_jwt_auth_config(tmp_path), workers=2 + ) as owned: + model: Final = scenario.model() + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(rs_token, options={"verify_signature": False})["sub"]) + ) + scenario.cleanups.callback( + scenario.delete_user, str(jwt.decode(ed_token, options={"verify_signature": False})["sub"]) + ) + + def chat(token: str, index: int) -> httpx.Response: + return owned.gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"d2-{index}"}]}, + key=token, + ) + + with ThreadPoolExecutor(max_workers=12) as pool: + futures = [pool.submit(chat, rs_token if index % 2 == 0 else ed_token, index) for index in range(24)] + owned_port: Final = owned.gateway.client.base_url.port + listeners: Final = [ + child + for child in psutil.Process(owned.process.pid).children(recursive=True) + if any( + connection.status == psutil.CONN_LISTEN and connection.laddr.port == owned_port + for connection in child.net_connections(kind="inet") + ) + ] + assert len(listeners) >= 2, listeners + killed_pid: Final = listeners[0].pid + listeners[0].kill() + results: Final[list[int]] = [] + for future in futures: + try: + result = future.result(timeout=60) + results.append(result.status_code) + except httpx.TransportError: + results.append(0) + assert all(status in (200, 401, 0) for status in results), results + assert any(status == 200 for status in results), results + dead: Final = eventually( + lambda: psutil.pid_exists(killed_pid), + lambda exists: not exists, + seconds=30, + ) + assert not dead, f"killed worker pid {killed_pid} still exists" + liveliness: Final = eventually( + lambda: owned.gateway.client.get("/health/liveliness"), + lambda response: response.status_code == 200, + seconds=45, + ) + assert liveliness.status_code == 200, liveliness.text diff --git a/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py new file mode 100644 index 00000000000..d00c51d1156 --- /dev/null +++ b/tests/integration/mcp/test_mcp_client_assertion_signing_alg.py @@ -0,0 +1,277 @@ +import json +import os +import uuid +from pathlib import Path +from typing import Final + +import httpx +import psycopg +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.mcp import forget_mcp, mcp_peer, register_mcp +from integration._support.process import owned_proxy_process + + +def _pem_private_key() -> str: + return ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + + +def _stored_client_assertion_signing_alg(identity: str) -> str: + rows: Final = read_rows('SELECT credentials FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + assert len(rows) == 1, rows + credentials: Final = rows[0]["credentials"] + if credentials is None: + return "RS256" + blob: Final = credentials if isinstance(credentials, dict) else json.loads(credentials) + return str(blob.get("client_assertion_signing_alg") or "RS256") + + +def test_non_approved_client_assertion_signing_alg_is_rejected_on_create_and_update(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "alg" + uuid.uuid4().hex[:8] + created: Final = gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": alias, + "alias": alias, + **peer.registration(), + "token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token", + "credentials": { + "client_private_key": _pem_private_key(), + "client_assertion_signing_alg": "HS256", + }, + }, + ) + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code in (400, 422), created.text + assert "client_assertion_signing_alg" in created.text, created.text + + identity: Final = register_mcp(scenario, peer, alias + "v") + edited: Final = gateway.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}}, + ) + assert edited.status_code in (400, 422), edited.text + assert "client_assertion_signing_alg" in edited.text, edited.text + assert _stored_client_assertion_signing_alg(identity) == "RS256" + + +def test_approved_client_assertion_signing_alg_round_trips(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "algok" + uuid.uuid4().hex[:8] + created: Final = gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": alias, + "alias": alias, + **peer.registration(), + "token_exchange_endpoint": "https://idp.integration.invalid/oauth2/token", + "credentials": { + "client_private_key": _pem_private_key(), + "client_assertion_signing_alg": "ES256", + }, + }, + ) + assert created.status_code == 201, created.text + identity: Final = str(created.json()["server_id"]) + scenario.cleanups.callback(forget_mcp, gateway, identity) + assert _stored_client_assertion_signing_alg(identity) == "ES256" + + +def test_server_row_with_stale_client_assertion_signing_alg_still_loads(gateway: Gateway) -> None: + """A row persisted before the allowlist (credentials.client_assertion_signing_alg = "HS256") + must keep loading with the RS256 fallback instead of disappearing on upgrade.""" + identity: Final = "stalealg" + uuid.uuid4().hex[:8] + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + 'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,' + " created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())", + ( + identity, + "stalealg" + uuid.uuid4().hex[:8], + "https://mcp.integration.invalid", + "http", + json.dumps({"client_assertion_signing_alg": "HS256"}), + ), + ) + try: + with mcp_peer() as peer, gateway.scenario() as scenario: + register_mcp(scenario, peer, "staletrigger" + uuid.uuid4().hex[:8]) + listed: Final = eventually( + lambda: gateway.request("GET", "/v1/mcp/server"), + lambda response: ( + response.status_code == 200 + and any(server.get("server_id") == identity for server in response.json()) + ), + seconds=30, + ) + assert listed.status_code == 200, listed.text + finally: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + + +APPROVED_ALGS: Final = ("RS256", "RS384", "RS512", "PS256", "PS384", "PS512", "ES256", "ES384", "ES512") + + +def _post_server(gateway: Gateway, name: str, credentials: object, omit_alg: bool = False) -> httpx.Response: + blob: Final = {} if omit_alg else {"client_assertion_signing_alg": credentials} + return gateway.request( + "POST", + "/v1/mcp/server", + { + "server_name": name, + "alias": name, + "url": "https://mcp.integration.invalid", + "transport": "http", + "credentials": blob, + }, + ) + + +def test_post_with_lowercase_hs256_is_rejected_naming_the_field(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b1hs" + uuid.uuid4().hex[:8], "hs256") + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code == 422, created.text + assert "client_assertion_signing_alg" in created.text, created.text + for algorithm in APPROVED_ALGS: + assert algorithm in created.text, created.text + + +def test_put_with_eddsa_on_existing_server_is_rejected(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "b2ed" + uuid.uuid4().hex[:8]) + edited: Final = gateway.request( + "PUT", + "/v1/mcp/server", + {"server_id": identity, "credentials": {"client_assertion_signing_alg": "EdDSA"}}, + ) + assert edited.status_code == 422, edited.text + assert "client_assertion_signing_alg" in edited.text, edited.text + + +@pytest.mark.parametrize("algorithm", APPROVED_ALGS) +def test_each_approved_client_assertion_signing_alg_is_accepted_and_stored(gateway: Gateway, algorithm: str) -> None: + with gateway.scenario() as scenario: + name: Final = "b3" + algorithm.lower() + uuid.uuid4().hex[:6] + created: Final = _post_server(gateway, name, algorithm) + assert created.status_code == 201, created.text + identity: Final = str(created.json()["server_id"]) + scenario.cleanups.callback(forget_mcp, gateway, identity) + listed: Final = gateway.request("GET", "/v1/mcp/server") + assert listed.status_code == 200, listed.text + assert any(server.get("server_id") == identity for server in listed.json()), listed.text + assert _stored_client_assertion_signing_alg(identity) == algorithm + + +@pytest.mark.parametrize( + ("label", "value", "expected"), + ( + ("empty", "", 422), + ("five-kb", "x" * 5120, 422), + ("integer", 7, 422), + ("list", ["RS256"], 422), + ("null", None, 201), + ), +) +def test_client_assertion_signing_alg_payload_variants( + gateway: Gateway, label: str, value: object, expected: int +) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, f"b{label}" + uuid.uuid4().hex[:8], value) + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + assert created.status_code == expected, created.text + if expected == 422 and label in ("empty", "five-kb"): + assert "client_assertion_signing_alg" in created.text, created.text + + +def test_server_post_without_credentials_alg_key_is_accepted(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b9absent" + uuid.uuid4().hex[:8], None, omit_alg=True) + assert created.status_code == 201, created.text + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + + +def test_unauthenticated_server_post_with_hs256_is_rejected(gateway: Gateway) -> None: + response: Final = gateway.client.post( + "/v1/mcp/server", + json={ + "server_name": "b10unauth" + uuid.uuid4().hex[:8], + "credentials": {"client_assertion_signing_alg": "hs256"}, + }, + ) + assert response.status_code == 401, response.text + + +def _seeded_row_loads_with_fallback_warning(gateway: Gateway, tmp_path: Path, algorithm: str) -> None: + identity: Final = "dualread" + uuid.uuid4().hex[:8] + name: Final = "dualread" + uuid.uuid4().hex[:8] + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + 'INSERT INTO "LiteLLM_MCPServerTable" (server_id, server_name, url, transport, credentials,' + " created_at, updated_at) VALUES (%s, %s, %s, %s, %s::jsonb, NOW(), NOW())", + ( + identity, + name, + "https://mcp.integration.invalid", + "http", + json.dumps({"client_assertion_signing_alg": algorithm}), + ), + ) + try: + with owned_proxy_process(gateway, tmp_path, {}) as owned: + listed: Final = eventually( + lambda: owned.gateway.request("GET", "/v1/mcp/server"), + lambda response: ( + response.status_code == 200 + and any(server.get("server_id") == identity for server in response.json()) + ), + seconds=30, + ) + assert listed.status_code == 200, listed.text + log_text: Final = eventually( + lambda: owned.log.read_text(), + lambda text: name in text and "not an approved algorithm" in text and "using RS256" in text, + seconds=30, + ) + assert name in log_text and "not an approved algorithm" in log_text and "using RS256" in log_text + finally: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute('DELETE FROM "LiteLLM_MCPServerTable" WHERE server_id = %s', (identity,)) + + +def test_seeded_hs256_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None: + _seeded_row_loads_with_fallback_warning(gateway, tmp_path, "HS256") + + +def test_seeded_eddsa_row_loads_with_rs256_fallback_warning(gateway: Gateway, tmp_path: Path) -> None: + _seeded_row_loads_with_fallback_warning(gateway, tmp_path, "EdDSA") + + +def test_chat_still_works_after_alg_rejection(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + created: Final = _post_server(gateway, "b13chat" + uuid.uuid4().hex[:8], "hs256") + assert created.status_code in (201, 422), created.text + if created.status_code == 201: + scenario.cleanups.callback(forget_mcp, gateway, str(created.json()["server_id"])) + model: Final = scenario.model() + body: Final = gateway.chat(model) + assert body["id"], body diff --git a/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py new file mode 100644 index 00000000000..4d75ba60d71 --- /dev/null +++ b/tests/integration/mcp/test_mcp_jwt_signer_jwks_allowlist.py @@ -0,0 +1,291 @@ +import base64 +import json +import os +import signal +import socket +import subprocess +import sys +import time +import uuid +from collections.abc import Callable, Generator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from pathlib import Path +from typing import Final + +import httpx +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import ed25519, rsa +from integration._support.client import Gateway, eventually +from integration._support.mcp import mcp_peer, register_mcp, tool_calls, tool_names +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import OKPAlgorithm, RSAAlgorithm + + +@contextmanager +def _jwt_signer_guardrail(gateway: Gateway, discovery_uri: str) -> Generator[None]: + created: Final = gateway.client.post( + "/guardrails", + headers={"x-litellm-api-key": gateway.key}, + json={ + "guardrail": { + "guardrail_name": "signer" + uuid.uuid4().hex[:8], + "litellm_params": { + "guardrail": "mcp_jwt_signer", + "mode": "pre_mcp_call", + "default_on": True, + "access_token_discovery_uri": discovery_uri, + }, + } + }, + ) + assert created.status_code == 200, created.text + identity: Final = str(created.json()["guardrail_id"]) + try: + yield + finally: + deleted: Final = gateway.client.delete(f"/guardrails/{identity}", headers={"x-litellm-api-key": gateway.key}) + assert deleted.status_code == 200, deleted.text + + +@contextmanager +def _idp_server(jwks_keys: list[Mapping[str, object]]) -> Generator[str]: + holder: Final = {"url": ""} + + def respond(request: Request) -> Reply: + if request.target == "/.well-known/openid-configuration": + return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode()) + if request.target == "/.well-known/jwks.json": + return Reply(body=json.dumps({"keys": list(jwks_keys)}).encode()) + return Reply(status=404) + + with wire_server(respond) as idp: + holder["url"] = idp.url + yield idp.url + "/.well-known/openid-configuration" + + +def _claims() -> dict[str, object]: + now: Final = int(time.time()) + return {"sub": "integration-mcp-user", "iat": now, "exp": now + 300} + + +def _hs256_key_and_token() -> tuple[dict[str, object], str]: + secret: Final = b"integration-hs256-client-secret-0123456789abcdef" + key: Final = { + "kty": "oct", + "alg": "HS256", + "kid": "sym", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + token: Final = jwt.encode(_claims(), secret, algorithm="HS256", headers={"kid": "sym"}) + return key, token + + +def _eddsa_key_and_token() -> tuple[dict[str, object], str]: + private_key: Final = ed25519.Ed25519PrivateKey.generate() + public_jwk: Final = json.loads(OKPAlgorithm.to_jwk(private_key.public_key())) + key: Final = {**public_jwk, "alg": "EdDSA", "kid": "ed"} + token: Final = jwt.encode(_claims(), private_key, algorithm="EdDSA", headers={"kid": "ed"}) + return key, token + + +def _rs256_key_and_token() -> tuple[dict[str, object], str]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) + key: Final = {**public_jwk, "alg": "RS256", "kid": "rsa"} + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + return key, token + + +def _call_with_bearer( + gateway: Gateway, key: str, identity: str, name: str, arguments: dict[str, object], bearer: str +) -> httpx.Response: + return gateway.client.post( + "/mcp-rest/tools/call", + headers={"x-litellm-api-key": key, "Authorization": f"Bearer {bearer}"}, + json={"server_id": identity, "name": name, "arguments": arguments}, + ) + + +@pytest.mark.parametrize("key_and_token", (_hs256_key_and_token, _eddsa_key_and_token), ids=("oct-HS256", "OKP-EdDSA")) +def test_jwks_key_with_non_approved_alg_cannot_verify_the_incoming_token( + gateway: Gateway, key_and_token: Callable[[], tuple[dict[str, object], str]] +) -> None: + jwks_key, token = key_and_token() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "jwksalg" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 401, response.text + assert "incoming token verification failed" in response.text, response.text + assert tool_calls(peer.drain()) == (), "rejected call reached the peer" + + +def test_jwks_rs256_key_still_verifies_the_incoming_token(gateway: Gateway) -> None: + jwks_key, token = _rs256_key_and_token() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "jwksok" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, peer, alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +def _rs256_keypair() -> tuple[rsa.RSAPrivateKey, dict[str, object]]: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return private_key, json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) + + +def _rs256_key_and_token_no_alg() -> tuple[dict[str, object], str]: + private_key, public_jwk = _rs256_keypair() + key: Final = {**public_jwk, "kid": "rsa"} + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + return key, token + + +def _signer_call(gateway: Gateway, key: str, identity: str, name: str, bearer: str) -> tuple[int, str]: + try: + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, bearer) + return response.status_code, response.text + except httpx.HTTPError as error: + return 0, repr(error) + + +def test_jwks_rs256_key_without_alg_field_still_verifies_the_token(gateway: Gateway) -> None: + jwks_key, token = _rs256_key_and_token_no_alg() + with _idp_server([jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksc4" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +def test_rs256_token_verifies_when_jwks_also_carries_a_non_approved_key(gateway: Gateway) -> None: + okp_key, _eddsa_token = _eddsa_key_and_token() + jwks_key, token = _rs256_key_and_token() + with _idp_server([okp_key, jwks_key]) as discovery_uri, _jwt_signer_guardrail(gateway, discovery_uri): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksc5" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + response: Final = _call_with_bearer(gateway, key, identity, name, {"a": 4, "b": 5}, token) + assert response.status_code == 200, response.text + assert response.json()["content"][0]["text"] == "9", response.text + assert len(tool_calls(peer.drain())) == 1, "accepted call never reached the peer" + + +@pytest.mark.parametrize( + ("label", "jwks_reply"), + ( + ("jwks-404", Reply(status=404, body=b'{"error": "missing"}')), + ("jwks-not-json", Reply(body=b"this is not json")), + ), + ids=("jwks-404", "jwks-not-json"), +) +def test_unusable_jwks_document_rejects_the_incoming_token(gateway: Gateway, label: str, jwks_reply: Reply) -> None: + _, token = _rs256_key_and_token() + + def respond(request: Request) -> Reply: + if request.target == "/.well-known/openid-configuration": + return Reply(body=json.dumps({"jwks_uri": holder["url"] + "/.well-known/jwks.json"}).encode()) + return jwks_reply + + holder: Final = {"url": ""} + with wire_server(respond) as idp: + holder["url"] = idp.url + with _jwt_signer_guardrail(gateway, idp.url + "/.well-known/openid-configuration"): + with mcp_peer() as peer, gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwks" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(gateway, key, identity)["add"] + peer.drain() + status, text = _signer_call(gateway, key, identity, name, token) + assert status == 401, (status, text) + assert "incoming token verification failed" in text, text + assert tool_calls(peer.drain()) == (), "rejected call reached the peer" + alive: Final = gateway.client.get("/health/liveliness") + assert alive.status_code == 200, alive.text + + +def test_signer_jwks_pause_rejects_then_recovers(gateway: Gateway, tmp_path: Path) -> None: + private_key, public_jwk = _rs256_keypair() + jwks_dir: Final = tmp_path / "jwks" / ".well-known" + jwks_dir.mkdir(parents=True) + port: Final = _reserve_port() + (jwks_dir / "jwks.json").write_text(json.dumps({"keys": [{**public_jwk, "kid": "rsa", "alg": "RS256"}]})) + (jwks_dir / "openid-configuration").write_text( + json.dumps({"jwks_uri": f"http://127.0.0.1:{port}/.well-known/jwks.json"}) + ) + token: Final = jwt.encode(_claims(), private_key, algorithm="RS256", headers={"kid": "rsa"}) + server: Final = subprocess.Popen( + [ + sys.executable, + "-I", + "-m", + "http.server", + str(port), + "--bind", + "127.0.0.1", + "--directory", + str(tmp_path / "jwks"), + ], + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + try: + with owned_proxy_process(gateway, tmp_path, {}) as owned: + discovery: Final = f"http://127.0.0.1:{port}/.well-known/openid-configuration" + with _jwt_signer_guardrail(owned.gateway, discovery): + with mcp_peer() as peer, owned.gateway.scenario() as scenario: + identity: Final = register_mcp(scenario, peer, "jwksd3" + uuid.uuid4().hex[:8]) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + name: Final = tool_names(owned.gateway, key, identity)["add"] + peer.drain() + os.kill(server.pid, signal.SIGSTOP) + outcomes: Final = [] + with ThreadPoolExecutor(max_workers=8) as pool: + futures = [ + pool.submit(_signer_call, owned.gateway, key, identity, name, token) for _ in range(8) + ] + for future in futures: + outcomes.append(future.result(timeout=80)) + assert all(status != 200 for status, _text in outcomes), outcomes + assert any(status in (0, 401, 500) for status, _text in outcomes), outcomes + os.kill(server.pid, signal.SIGCONT) + recovered: Final = eventually( + lambda: _signer_call(owned.gateway, key, identity, name, token), + lambda outcome: outcome[0] == 200, + seconds=70, + ) + assert recovered[0] == 200, recovered + assert tool_calls(peer.drain()) != (), "resumed call never reached the peer" + finally: + try: + os.kill(server.pid, signal.SIGCONT) + except ProcessLookupError: + pass + server.terminate() + server.wait(timeout=10) + + +def _reserve_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py index db3a1a386a3..29d1f6b693b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_token_endpoint.py @@ -16,7 +16,7 @@ import jwt import litellm import pytest from cryptography.hazmat.primitives import serialization -from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.hazmat.primitives.asymmetric import ec, rsa from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( Error, @@ -50,6 +50,15 @@ _PRIVATE_PEM = _RSA_KEY.private_bytes( serialization.PrivateFormat.PKCS8, serialization.NoEncryption(), ).decode() +_EC_PRIVATE_PEM = ( + ec.generate_private_key(ec.SECP256R1()) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() +) _PUBLIC_PEM = ( _RSA_KEY.public_key() .public_bytes( @@ -185,9 +194,9 @@ async def test_fetch_network_error_maps_to_upstream_unavailable(raised): "auth", [ PrivateKeyJwtAuth(private_key=SecretStr("not-a-pem-key"), signing_alg="RS256"), - PrivateKeyJwtAuth(private_key=SecretStr(_PRIVATE_PEM), signing_alg="XX999"), + PrivateKeyJwtAuth(private_key=SecretStr(_EC_PRIVATE_PEM), signing_alg="RS256"), ], - ids=["garbage-key", "unknown-alg"], + ids=["garbage-key", "key-alg-mismatch"], ) async def test_fetch_unsignable_client_assertion_is_misconfigured_not_a_crash(auth): client = AsyncMock() diff --git a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py index 2e255ddf853..bf16902898d 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_db_credentials.py @@ -1794,3 +1794,20 @@ async def test_unverified_legacy_cache_cannot_bypass_enforcement(monkeypatch): await mcp_per_user_token_cache.set("alice", "srv", "bob", 60) assert await module.resolve_user_oauth_access_token("alice", server) is None assert await mcp_per_user_token_cache.get("alice", "srv") is None + + +def test_server_table_row_with_stale_client_assertion_signing_alg_still_validates() -> None: + """Rows written before the approved-algorithm allowlist carry values like HS256 in the + credentials blob; the stored shape stays lenient so reload_servers_from_database still + parses the row and _stored_client_assertion_signing_alg falls back to RS256 instead of + the server silently disappearing on upgrade.""" + row: Final = { + "server_id": "srv-stale-alg", + "server_name": "stale_alg_server", + "url": "https://mcp.example.com", + "transport": "http", + "credentials": {"client_assertion_signing_alg": "HS256"}, + } + parsed: Final = LiteLLM_MCPServerTable.model_validate(row) + assert parsed.credentials is not None + assert parsed.credentials["client_assertion_signing_alg"] == "HS256" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 123a5505953..bf162027f14 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -1,7 +1,9 @@ -from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth -import importlib import asyncio + +# Add the parent directory to the path so we can import litellm +import contextlib import functools +import importlib import json import logging import os @@ -10,26 +12,13 @@ import time from collections.abc import AsyncIterator from datetime import datetime from pathlib import Path -from typing import Any, Dict, Final, Literal, Optional +from typing import Any, Final, Literal from unittest.mock import AsyncMock, MagicMock, patch -import pytest -from fastapi import HTTPException -from respx import MockRouter - -from litellm.proxy._experimental.mcp_server.exceptions import ( - MCPServerListError, - MCPUpstreamAuthError, -) -from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListFault - -# Add the parent directory to the path so we can import litellm - - -import contextlib - import httpx import httpx2 +import pytest +from fastapi import HTTPException from mcp import ReadResourceResult, Resource from mcp.types import ( CallToolResult, @@ -40,24 +29,37 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, TypeAdapter +from respx import MockRouter +import litellm +import litellm.llms as litellm_llms +from litellm.caching.caching import DualCache +from litellm.caching.llm_caching_handler import LLMClientCache from litellm.constants import MCP_METADATA_TIMEOUT +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._experimental.mcp_server import discoverable_endpoints -from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult +from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPServerListError, + MCPUpstreamAuthError, +) +from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ServerListFault from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( ListedToolsCaller, MCPServerManager, _deserialize_json_dict, + _deserialize_json_list, _flow_endpoints_missing, _mcp_oauth_discovery_on_startup_enabled, - _oauth_endpoints_unresolved, - _deserialize_json_list, _normalize_mcp_server_cost_info, + _oauth_endpoints_unresolved, _obo_retry_applies, _resolve_openapi_tool_auth, _should_strip_caller_authorization, listed_tools_caller_for, ) +from litellm.proxy._experimental.mcp_server.tool_outcome import TextResult +from litellm.proxy._experimental.mcp_server.upstream import resolve_upstream_auth from litellm.proxy._types import ( LiteLLM_MCPServerTable, LiteLLM_ObjectPermissionTable, @@ -68,18 +70,12 @@ from litellm.proxy._types import ( MCPTransport, UserAPIKeyAuth, ) -from litellm.types.llms.custom_http import httpxSpecialProvider -from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol -from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool -from litellm.caching.caching import DualCache -from litellm.caching.llm_caching_handler import LLMClientCache -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -import litellm -from litellm.integrations.custom_guardrail import CustomGuardrail -import litellm.llms as litellm_llms from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks from litellm.types.integrations.slack_alerting import AlertType +from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.types.mcp import MCPAuth, MCPAuthType, MCPUpstreamProtocol +from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer, PinnedMCPTool @pytest.mark.asyncio @@ -2134,7 +2130,7 @@ class TestMCPServerManager: assert exc_info.value.fault == ServerListFault(tag="internal", status_code=412) assert exc_info.value.server_name == "te-412-server" - def _upstream_status_error(self, status_code: int, www_authenticate: Optional[str] = None) -> httpx.HTTPStatusError: + def _upstream_status_error(self, status_code: int, www_authenticate: str | None = None) -> httpx.HTTPStatusError: """Build an httpx.HTTPStatusError shaped like the one the MCP SDK surfaces for an upstream HTTP failure, so _extract_upstream_auth_failure can read status_code and WWW-Authenticate.""" request = httpx.Request("POST", "https://up.example.com/mcp") @@ -2804,11 +2800,11 @@ class TestMCPServerManager: assert built.token_url == "https://idp.example.com/manual-token" assert built.scopes == ["calendar.read"] - async def _capture_subject_token(self, call) -> Optional[str]: + async def _capture_subject_token(self, call) -> str | None: """Run a manager method (via ``call(manager)``) and return the subject_token it threaded into ``_create_mcp_client``.""" manager = MCPServerManager() - captured: Dict[str, Any] = {} + captured: dict[str, Any] = {} async def capture_create_mcp_client( server, mcp_auth_header, extra_headers, stdio_env, subject_token=None, **kwargs @@ -9103,7 +9099,7 @@ class TestMCPServerTimestamps: _carry_forward_resolved_oauth_endpoints, ) - def make_server(url: str, auth_type: MCPAuth, authorization_url: Optional[str]) -> MCPServer: + def make_server(url: str, auth_type: MCPAuth, authorization_url: str | None) -> MCPServer: return MCPServer( server_id="s1", name="s1", @@ -10476,7 +10472,7 @@ class TestOAuthDiscoverySSRFGuard: } mock_response.raise_for_status = MagicMock() - captured_kwargs: Dict[str, Any] = {} + captured_kwargs: dict[str, Any] = {} async def fake_get(url, **kwargs): captured_kwargs.update(kwargs) @@ -10507,7 +10503,7 @@ class TestOAuthDiscoverySSRFGuard: } mock_response.raise_for_status = MagicMock() - captured_kwargs: Dict[str, Any] = {} + captured_kwargs: dict[str, Any] = {} async def fake_get(url, **kwargs): captured_kwargs.update(kwargs) @@ -10748,8 +10744,8 @@ class TestHealthCheckInterpolatesGlobalEnvVars: ) @staticmethod - def _capture_headers(manager: MCPServerManager) -> Dict[str, Any]: - captured: Dict[str, Any] = {} + def _capture_headers(manager: MCPServerManager) -> dict[str, Any]: + captured: dict[str, Any] = {} mock_client = AsyncMock() mock_client.run_with_session = AsyncMock(return_value="ok") @@ -10796,7 +10792,7 @@ class TestUserEnvVarsCacheEviction: def _patch_cache(monkeypatch, max_size): from litellm.proxy._experimental.mcp_server import mcp_server_manager as m - cache: Dict[Any, Any] = {} + cache: dict[Any, Any] = {} monkeypatch.setattr(m, "_user_env_vars_cache", cache) monkeypatch.setattr(m, "_USER_ENV_VARS_CACHE_MAX_SIZE", max_size) return m, cache @@ -11022,7 +11018,7 @@ class TestCreateMcpClientV2Graft: """ def _http_server(self, **overrides: Any) -> MCPServer: - base: Dict[str, Any] = dict( + base: dict[str, Any] = dict( server_id="http-graft", name="graft_server", url="https://upstream.example.com/mcp", @@ -13943,7 +13939,7 @@ class TestConfigServerIdPinning: @pytest.mark.asyncio async def test_two_servers_pinning_the_same_id_are_rejected(self, config_only_mcp_manager_factory): manager = config_only_mcp_manager_factory() - config: Dict[str, Any] = { + config: dict[str, Any] = { "docs_server": {"url": "https://a.example.com/mcp", "server_id": "shared-id"}, "wiki_server": {"url": "https://b.example.com/mcp", "server_id": "shared-id"}, } @@ -13962,7 +13958,7 @@ class TestConfigServerIdPinning: auth_type=None, alias=None, ) - config: Dict[str, Any] = { + config: dict[str, Any] = { "docs_server": {"url": "https://a.example.com/mcp", "transport": MCPTransport.http}, "wiki_server": {"url": "https://b.example.com/mcp", "server_id": derived}, } @@ -14793,11 +14789,11 @@ async def test_debug_resolution_matches_final_header_conflict_winner( expected_source: str, expected_authorization: str | None, ) -> None: - from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var - from starlette.requests import Request from pydantic import SecretStr + from starlette.requests import Request from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import MCPAuthenticatedUser + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from litellm.proxy._experimental.mcp_server.mcp_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics from litellm.proxy._experimental.mcp_server.outbound_credentials import ( ApiKeyConfig, @@ -14863,9 +14859,9 @@ async def test_debug_reports_legacy_signing_and_non_http_transport( _mcp_request_ctx, monkeypatch, transport: Literal["http", "stdio"] ) -> None: monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") - from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from starlette.requests import Request + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var from litellm.proxy._experimental.mcp_server.mcp_debug import MCP_AUTH_DIAGNOSTICS_SCOPE_KEY, MCPAuthDiagnostics from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -15172,7 +15168,6 @@ class _DiscoveryClock: return self.now -from pydantic import TypeAdapter from mcp.types import JSONRPCMessage _JSONRPC_ADAPTER = TypeAdapter(JSONRPCMessage) @@ -15269,7 +15264,6 @@ def _discovery_server() -> MCPServer: @pytest.mark.asyncio @pytest.mark.parametrize("kind", ("prompts", "resources", "templates")) async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None: - import respx clock: Final = _DiscoveryClock() manager: Final = MCPServerManager(discovery_clock=clock) @@ -15301,7 +15295,6 @@ async def test_discovery_cache_reuses_raw_results_and_expires(kind: str) -> None @pytest.mark.parametrize("kind", ("prompts", "resources", "templates")) @pytest.mark.parametrize("outcome", ("unsupported", "rejected", "failure")) async def test_discovery_cache_empty_results_and_failures(kind: str, outcome: str) -> None: - import respx manager: Final = MCPServerManager() upstream: Final = _DiscoveryUpstream() @@ -15346,7 +15339,6 @@ async def test_discovery_cache_retries_failed_pagination_before_caching_complete @pytest.mark.asyncio async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_auth() -> None: - import respx manager: Final = MCPServerManager() upstream: Final = _DiscoveryUpstream() @@ -15376,7 +15368,6 @@ async def test_discovery_cache_isolates_forwarded_credentials_and_shares_static_ @pytest.mark.asyncio async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> None: - import respx manager: Final = MCPServerManager() upstream: Final = _DiscoveryUpstream() @@ -15400,7 +15391,6 @@ async def test_discovery_cache_coalesces_and_survives_waiter_cancellation() -> N @pytest.mark.asyncio async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old_results() -> None: - import respx manager: Final = MCPServerManager() upstream: Final = _DiscoveryUpstream() @@ -15420,7 +15410,6 @@ async def test_discovery_cache_invalidation_during_fetch_does_not_repopulate_old @pytest.mark.asyncio async def test_discovery_cache_can_be_disabled(monkeypatch: pytest.MonkeyPatch) -> None: - import respx monkeypatch.setenv("LITELLM_MCP_DISCOVERY_CACHE_TTL", "0") manager: Final = MCPServerManager() @@ -15586,7 +15575,6 @@ async def test_discovery_cache_bounds_detached_fetches_without_dropping_results( @pytest.mark.asyncio async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> None: - import respx from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import StaticHeaderAuth from litellm.proxy._experimental.mcp_server.outbound_credentials.resolver import UpstreamCredentialProvider from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok, Result @@ -15649,7 +15637,6 @@ async def test_discovery_cache_tracks_resolved_credentials_across_workers() -> N @pytest.mark.asyncio async def test_discovery_resolves_stored_oauth_for_the_requesting_user() -> None: - import respx from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import OAuthToken class TokenStore: @@ -16090,6 +16077,7 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio async def test_static_resolution_cancellation_closes_flow(self) -> None: from collections.abc import AsyncGenerator + from litellm.experimental_mcp_client.client import MCPClient from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import prepare_mcp_client @@ -16355,8 +16343,8 @@ class TestProtectedCredentialPreparation: @pytest.mark.asyncio @pytest.mark.parametrize("selected", [False, True]) async def test_request_selected_during_guardrail_runs_concurrently_with_tool(monkeypatch, selected): - from litellm.responses.mcp.request_context import MCPRequestContext from litellm.proxy._experimental.mcp_server import tool_registry + from litellm.responses.mcp.request_context import MCPRequestContext tool_started = asyncio.Event() guardrail_started = asyncio.Event() @@ -16434,6 +16422,7 @@ async def test_server_response_identifies_read_only_config(in_config, in_db, exp @pytest.mark.parametrize("with_caller,legacy_factory", [(True, False), (False, False), (True, True)]) async def test_client_sampling_does_not_fill_explicit_context_from_another_ambient_caller(with_caller, legacy_factory): from mcp.server.auth.middleware.auth_context import auth_context_var + from litellm.proxy._experimental.mcp_server import server as legacy_server from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback @@ -16546,6 +16535,40 @@ class TestSharedIdentifierPrefixWarning: assert "'shared'" in shared_warnings[0] +def test_stored_client_assertion_signing_alg_falls_back_to_rs256_for_non_approved(caplog): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _stored_client_assertion_signing_alg, + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert _stored_client_assertion_signing_alg("HS256", "srv") == "RS256" + assert "client_assertion_signing_alg" in caplog.text and "HS256" in caplog.text and "srv" in caplog.text + + +def test_stored_client_assertion_signing_alg_passes_through_approved(caplog): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _stored_client_assertion_signing_alg, + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + assert _stored_client_assertion_signing_alg("PS384", "srv") == "PS384" + assert _stored_client_assertion_signing_alg(None, "srv") == "RS256" + assert "approved algorithm" not in caplog.text + + +def test_mcp_server_model_rejects_non_approved_client_assertion_signing_alg(): + from pydantic import ValidationError + + with pytest.raises(ValidationError) as exc: + MCPServer( + server_id="srv", + name="srv", + transport="http", + client_assertion_signing_alg="EdDSA", + ) + assert "client_assertion_signing_alg" in str(exc.value) + + @pytest.mark.asyncio @pytest.mark.parametrize( "flag,transports,expected_warnings", @@ -17244,8 +17267,11 @@ async def test_upstream_preparation_honors_case_sensitive_extra_command(monkeypa monkeypatch.setenv("LITELLM_ENABLE_MCP_STDIO", "true") monkeypatch.setattr(upstream, "MCP_STDIO_ALLOWED_COMMANDS", frozenset({"CustomRunner"})) server: Final = MCPServer( - server_id="custom-stdio", name="custom-stdio", transport=MCPTransport.stdio, - command="/opt/tools/CustomRunner", args=[], + server_id="custom-stdio", + name="custom-stdio", + transport=MCPTransport.stdio, + command="/opt/tools/CustomRunner", + args=[], ) client: Final = await MCPServerManager()._create_mcp_client(server) assert client.stdio_config is not None diff --git a/tests/unit/proxy/_experimental/mcp_server/test_oauth_identity_binding.py b/tests/unit/proxy/_experimental/mcp_server/test_oauth_identity_binding.py index 0036035f448..ee1623fe22a 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_oauth_identity_binding.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_oauth_identity_binding.py @@ -665,3 +665,35 @@ async def test_audit_matching_login_without_nonce_does_not_report_failure(caplog ) assert result is None assert "oauth_identity_binding audit" not in caplog.text + + +def test_select_signing_key_rejects_kid_matching_key_with_non_approved_alg() -> None: + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + from litellm.proxy._experimental.mcp_server.oauth_identity_binding import _BindingRejection + + okp_jwk: Final = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())), + "kid": KID, + "alg": "EdDSA", + } + result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk]) + assert isinstance(result, _BindingRejection) + + +def test_select_signing_key_picks_approved_key_when_kid_is_shared() -> None: + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + okp_jwk: Final = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(ed25519.Ed25519PrivateKey.generate().public_key())), + "kid": KID, + "alg": "EdDSA", + } + hs_jwk: Final = {"kty": "oct", "kid": KID, "alg": "HS256", "k": "c2VjcmV0"} + result: Final = _select_signing_key(_sign_id_token({}), [okp_jwk, hs_jwk, _PUBLIC_JWK]) + assert isinstance(result, jwt.PyJWK) + assert result.key_id == KID diff --git a/tests/unit/proxy/auth/test_handle_jwt.py b/tests/unit/proxy/auth/test_handle_jwt.py index 640b3d8053d..4f51a24713c 100644 --- a/tests/unit/proxy/auth/test_handle_jwt.py +++ b/tests/unit/proxy/auth/test_handle_jwt.py @@ -104,9 +104,7 @@ async def test_map_user_to_teams_handles_already_in_team_exception(): ) as mock_add: with patch("litellm.proxy.auth.handle_jwt.verbose_proxy_logger") as mock_logger: # This should not raise an exception - result = await JWTAuthManager.map_user_to_teams( - user_object=user, team_object=team - ) + result = await JWTAuthManager.map_user_to_teams(user_object=user, team_object=team) # Verify the method completed successfully assert result is None @@ -143,14 +141,10 @@ async def test_map_user_to_teams_reraises_other_proxy_exceptions(): async def test_map_user_to_teams_null_inputs(): """Test that method handles null inputs gracefully""" # Test with null user - await JWTAuthManager.map_user_to_teams( - user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1") - ) + await JWTAuthManager.map_user_to_teams(user_object=None, team_object=LiteLLM_TeamTable(team_id="test_team_1")) # Test with null team - await JWTAuthManager.map_user_to_teams( - user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None - ) + await JWTAuthManager.map_user_to_teams(user_object=LiteLLM_UserTable(user_id="test_user_1"), team_object=None) # Test with both null await JWTAuthManager.map_user_to_teams(user_object=None, team_object=None) @@ -206,9 +200,7 @@ async def test_find_team_with_model_access_reports_passthrough_allowlist_denial( assert exc_info.value.status_code == 403 assert "allowed_passthrough_routes" in exc_info.value.detail assert "requested model" not in exc_info.value.detail - mock_is_auth_enforced_pass_through_route.assert_called_once_with( - route="/my-pass-through", method="POST" - ) + mock_is_auth_enforced_pass_through_route.assert_called_once_with(route="/my-pass-through", method="POST") user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] assert user_api_key_dict.metadata == {} @@ -401,9 +393,7 @@ async def test_auth_builder_proxy_admin_user_role(): route = "/chat/completions" # Create user object with PROXY_ADMIN role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.PROXY_ADMIN) # Create mock JWT handler jwt_handler = JWTHandler() @@ -412,14 +402,10 @@ async def test_auth_builder_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -427,9 +413,7 @@ async def test_auth_builder_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -442,9 +426,7 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -457,12 +439,8 @@ async def test_auth_builder_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -496,9 +474,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): route = "/chat/completions" # Create user object with regular USER role - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) # Create mock JWT handler jwt_handler = JWTHandler() @@ -507,14 +483,10 @@ async def test_auth_builder_non_proxy_admin_user_role(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -522,9 +494,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): return_value=("test_user_1", "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -537,9 +507,7 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(None, None), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -552,12 +520,8 @@ async def test_auth_builder_non_proxy_admin_user_role(): new_callable=AsyncMock, return_value=(user_object, None, None, None, user_object.user_id), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} @@ -706,11 +670,7 @@ async def test_sync_user_role_and_teams(): prisma_client=None, user_api_key_cache=mock_user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -719,9 +679,7 @@ async def test_sync_user_role_and_teams(): token = {"roles": ["ADMIN"], "my_id_teams": ["team1", "team2"]} - user = LiteLLM_UserTable( - user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"] - ) + user = LiteLLM_UserTable(user_id="u1", user_role=LitellmUserRoles.INTERNAL_USER.value, teams=["team2"]) prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() @@ -748,11 +706,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -769,9 +723,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_role_change(): prisma = AsyncMock() prisma.db.litellm_usertable.update = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -791,11 +743,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -816,9 +764,7 @@ async def test_sync_user_role_and_teams_cache_invalidation_on_team_change(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ): - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_called_once() call_kwargs = mock_cache.async_set_cache.call_args @@ -838,11 +784,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma_client=None, user_api_key_cache=AsyncMock(), litellm_jwtauth=LiteLLM_JWTAuth( - jwt_litellm_role_map=[ - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ) - ], + jwt_litellm_role_map=[JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN)], roles_jwt_field="roles", team_ids_jwt_field="my_id_teams", sync_user_role_and_teams=True, @@ -858,9 +800,7 @@ async def test_sync_user_role_and_teams_no_cache_write_when_nothing_changes(): prisma = AsyncMock() - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, prisma, user_api_key_cache=mock_cache - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, prisma, user_api_key_cache=mock_cache) mock_cache.async_set_cache.assert_not_called() @@ -885,9 +825,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): assert jwt_handler.get_all_jwt_team_ids({"teams": ["a", "b"]}) == ["a", "b"] # both populated, no overlap - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": "primary", "teams": ["a", "b"]} - ) == ["a", "b", "primary"] + assert jwt_handler.get_all_jwt_team_ids({"team_id": "primary", "teams": ["a", "b"]}) == ["a", "b", "primary"] # both populated with overlap — singular dedup'd assert jwt_handler.get_all_jwt_team_ids({"team_id": "a", "teams": ["a", "b"]}) == [ @@ -896,9 +834,7 @@ def test_get_all_jwt_team_ids_unions_singular_and_plural(): ] # singular field as multi-element list (some IdPs) — merge all, preserve plural-first order - assert jwt_handler.get_all_jwt_team_ids( - {"team_id": ["primary", "secondary"], "teams": ["a"]} - ) == [ + assert jwt_handler.get_all_jwt_team_ids({"team_id": ["primary", "secondary"], "teams": ["a"]}) == [ "a", "primary", "secondary", @@ -958,19 +894,11 @@ async def test_map_jwt_role_to_litellm_role(): litellm_jwtauth=LiteLLM_JWTAuth( jwt_litellm_role_map=[ # Exact match - JWTLiteLLMRoleMap( - jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN - ), + JWTLiteLLMRoleMap(jwt_role="ADMIN", litellm_role=LitellmUserRoles.PROXY_ADMIN), # Wildcard patterns - JWTLiteLLMRoleMap( - jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER - ), - JWTLiteLLMRoleMap( - jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM - ), - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="user_*", litellm_role=LitellmUserRoles.INTERNAL_USER), + JWTLiteLLMRoleMap(jwt_role="team_?", litellm_role=LitellmUserRoles.TEAM), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ], roles_jwt_field="roles", ), @@ -1042,9 +970,7 @@ async def test_map_jwt_role_to_litellm_role(): # Test patterns that don't match character classes jwt_handler.litellm_jwtauth.jwt_litellm_role_map = [ - JWTLiteLLMRoleMap( - jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER - ), + JWTLiteLLMRoleMap(jwt_role="dev_[123]", litellm_role=LitellmUserRoles.INTERNAL_USER), ] token = {"roles": ["dev_4"]} # 4 is not in [123] result = jwt_handler.map_jwt_role_to_litellm_role(token) @@ -1139,25 +1065,19 @@ async def test_nested_jwt_field_access(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(nested_token, None) == "obj789" # Test 5b: object_id_jwt_field with flat access (backward compatibility) jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(flat_token, None) == "obj789" # Test 6: end_user_id_jwt_field with nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") assert jwt_handler.get_end_user_id(nested_token, None) == "customer123" # Test 6b: end_user_id_jwt_field with flat access (backward compatibility) @@ -1173,9 +1093,7 @@ async def test_nested_jwt_field_access(): assert jwt_handler.get_team_id(flat_token, None) == "team456" # Test 8: roles_jwt_field with deeply nested access (already supported, but testing) - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") assert jwt_handler.get_jwt_role(nested_token, []) == ["admin", "user"] # Test 9: user_roles_jwt_field with nested access (already supported, but testing) @@ -1221,10 +1139,7 @@ async def test_nested_jwt_field_missing_paths(): # Test 2: Missing user.email should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="user.email") - assert ( - jwt_handler.get_user_email(incomplete_token, "default@example.com") - == "default@example.com" - ) + assert jwt_handler.get_user_email(incomplete_token, "default@example.com") == "default@example.com" # Test 3: Missing groups should return empty list jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") @@ -1239,43 +1154,28 @@ async def test_nested_jwt_field_missing_paths(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( object_id_jwt_field="profile.object_id", - role_mappings=[ - RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER) - ], + role_mappings=[RoleMapping(role="admin", internal_role=LitellmUserRoles.INTERNAL_USER)], ) assert jwt_handler.get_object_id(incomplete_token, "default_obj") == "default_obj" # Test 6: Missing customer.end_user_id should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - end_user_id_jwt_field="customer.end_user_id" - ) - assert ( - jwt_handler.get_end_user_id(incomplete_token, "default_customer") - == "default_customer" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(end_user_id_jwt_field="customer.end_user_id") + assert jwt_handler.get_end_user_id(incomplete_token, "default_customer") == "default_customer" # Test 7: Missing tenant.team_id should use team_id_default fallback - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="tenant.team_id", team_id_default="fallback_team" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="tenant.team_id", team_id_default="fallback_team") assert jwt_handler.get_team_id(incomplete_token, "default_team") == "fallback_team" # Test 8: Missing resource_access.my-client.roles should return default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - roles_jwt_field="resource_access.my-client.roles" - ) - assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == [ - "default_role" - ] + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(roles_jwt_field="resource_access.my-client.roles") + assert jwt_handler.get_jwt_role(incomplete_token, ["default_role"]) == ["default_role"] # Test 9: Missing nested user roles should return default jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( user_roles_jwt_field="resource_access.my-client.roles", user_allowed_roles=["admin", "user"], ) - assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == [ - "default_user_role" - ] + assert jwt_handler.get_user_roles(incomplete_token, ["default_user_role"]) == ["default_user_role"] @pytest.mark.asyncio @@ -1300,9 +1200,7 @@ async def test_metadata_prefix_handling_in_nested_fields(): } # Test 1: metadata.user.email should access user.email after prefix removal - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - user_email_jwt_field="metadata.user.email" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(user_email_jwt_field="metadata.user.email") # The get_nested_value function removes "metadata." prefix, so "metadata.user.email" becomes "user.email" assert jwt_handler.get_user_email(token, None) == "user@example.com" @@ -1338,9 +1236,7 @@ async def test_find_team_with_model_access_model_group(monkeypatch): async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1448,9 +1344,7 @@ async def test_find_team_with_model_access_v1_messages_default_routes(monkeypatc async def mock_get_team_object(*args, **kwargs): # type: ignore return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -1516,14 +1410,10 @@ async def test_auth_builder_returns_team_membership_object(): team_id=_team_id, budget_id="budget_123", spend=10.5, - litellm_budget_table=LiteLLM_BudgetTable( - budget_id="budget_123", rpm_limit=100, tpm_limit=5000 - ), + litellm_budget_table=LiteLLM_BudgetTable(budget_id="budget_123", rpm_limit=100, tpm_limit=5000), ) - user_object = LiteLLM_UserTable( - user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id=_user_id, user_role=LitellmUserRoles.INTERNAL_USER) team_object = LiteLLM_TeamTable(team_id=_team_id) @@ -1534,14 +1424,10 @@ async def test_auth_builder_returns_team_membership_object(): # Mock all the dependencies and method calls with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, - patch.object( - JWTAuthManager, "check_rbac_role", new_callable=AsyncMock - ) as mock_check_rbac, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock) as mock_check_rbac, patch.object(jwt_handler, "get_rbac_role", return_value=None) as mock_get_rbac, patch.object(jwt_handler, "get_scopes", return_value=[]) as mock_get_scopes, - patch.object( - jwt_handler, "get_object_id", return_value=None - ) as mock_get_object_id, + patch.object(jwt_handler, "get_object_id", return_value=None) as mock_get_object_id, patch.object( JWTAuthManager, "get_user_info", @@ -1549,9 +1435,7 @@ async def test_auth_builder_returns_team_membership_object(): return_value=(_user_id, "test@example.com", True), ) as mock_get_user_info, patch.object(jwt_handler, "get_org_id", return_value=None) as mock_get_org_id, - patch.object( - jwt_handler, "get_end_user_id", return_value=None - ) as mock_get_end_user_id, + patch.object(jwt_handler, "get_end_user_id", return_value=None) as mock_get_end_user_id, patch.object( JWTAuthManager, "check_admin_access", @@ -1564,9 +1448,7 @@ async def test_auth_builder_returns_team_membership_object(): new_callable=AsyncMock, return_value=(_team_id, team_object), ) as mock_find_team, - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=set() - ) as mock_get_all_team_ids, + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()) as mock_get_all_team_ids, patch.object( JWTAuthManager, "find_team_with_model_access", @@ -1585,15 +1467,9 @@ async def test_auth_builder_returns_team_membership_object(): user_object.user_id, ), ) as mock_get_objects, - patch.object( - JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock - ) as mock_map_user, - patch.object( - JWTAuthManager, "validate_object_id", return_value=True - ) as mock_validate_object, - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ) as mock_sync_user, + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock) as mock_map_user, + patch.object(JWTAuthManager, "validate_object_id", return_value=True) as mock_validate_object, + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock) as mock_sync_user, ): # Set up the mock return values mock_auth_jwt.return_value = {"sub": _user_id, "scope": ""} @@ -1612,24 +1488,12 @@ async def test_auth_builder_returns_team_membership_object(): ) # Verify that team_membership_object is returned - assert result["team_membership"] is not None, ( - "team_membership should be present" - ) - assert result["team_membership"] == mock_team_membership, ( - "team_membership should match the mock object" - ) - assert result["team_membership"].user_id == _user_id, ( - "team_membership user_id should match" - ) - assert result["team_membership"].team_id == _team_id, ( - "team_membership team_id should match" - ) - assert result["team_membership"].budget_id == "budget_123", ( - "team_membership budget_id should match" - ) - assert result["team_membership"].spend == 10.5, ( - "team_membership spend should match" - ) + assert result["team_membership"] is not None, "team_membership should be present" + assert result["team_membership"] == mock_team_membership, "team_membership should match the mock object" + assert result["team_membership"].user_id == _user_id, "team_membership user_id should match" + assert result["team_membership"].team_id == _team_id, "team_membership team_id should match" + assert result["team_membership"].budget_id == "budget_123", "team_membership budget_id should match" + assert result["team_membership"].spend == 10.5, "team_membership spend should match" @pytest.mark.asyncio @@ -1864,9 +1728,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() general_settings = {"enforce_rbac": False} route = "/chat/completions" - user_object = LiteLLM_UserTable( - user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="test_user_1", user_role=LitellmUserRoles.INTERNAL_USER) jwt_handler = JWTHandler() user_api_key_cache = DualCache() @@ -1885,9 +1747,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() jwt_response = {"sub": "test_user_1", "scope": ""} with ( - patch.object( - jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock - ) as mock_get_userinfo, + patch.object(jwt_handler, "get_oidc_userinfo", new_callable=AsyncMock) as mock_get_userinfo, patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), @@ -1928,9 +1788,7 @@ async def test_auth_builder_oidc_enabled_falls_back_to_jwt_auth_for_jwt_tokens() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = jwt_response @@ -2012,9 +1870,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): ) team_object = LiteLLM_TeamTable(team_id="team-2") - user_object = LiteLLM_UserTable( - user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER - ) + user_object = LiteLLM_UserTable(user_id="user-1", user_role=LitellmUserRoles.INTERNAL_USER) with ( patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, @@ -2025,9 +1881,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): new_callable=AsyncMock, return_value=None, ), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, patch.object( JWTAuthManager, "get_objects", @@ -2035,9 +1889,7 @@ async def test_auth_builder_uses_team_from_header_e2e(): return_value=(user_object, None, None, None, user_object.user_id), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), ): mock_auth_jwt.return_value = { "sub": "user-1", @@ -2254,9 +2106,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): return_value=(None, None, None, None, None), ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.RouteChecks.is_auth_enforced_pass_through_route", return_value=True, @@ -2285,9 +2135,7 @@ async def test_auth_builder_rbac_team_loads_team_for_passthrough_allowlist(): mock_get_team.assert_awaited_once() assert mock_get_team.await_args.kwargs["team_id"] == "team-rbac" user_api_key_dict = mock_passthrough_check.call_args.kwargs["user_api_key_dict"] - assert user_api_key_dict.team_metadata == { - "allowed_passthrough_routes": ["/my-pass-through"] - } + assert user_api_key_dict.team_metadata == {"allowed_passthrough_routes": ["/my-pass-through"]} @pytest.mark.asyncio @@ -2381,9 +2229,7 @@ async def test_auth_builder_admin_on_llm_route_honors_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2435,9 +2281,7 @@ async def test_auth_builder_admin_on_mgmt_route_ignores_team_header(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2491,9 +2335,7 @@ async def test_auth_builder_admin_on_llm_route_without_header_unchanged(): patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "is_admin", return_value=True), - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_team, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_team, ): mock_auth_jwt.return_value = { "sub": "admin-user", @@ -2537,9 +2379,7 @@ async def test_get_team_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="organization.team.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="organization.team.name") assert jwt_handler.get_team_alias(nested_token, None) == "engineering-team" # Test flat access (backward compatibility) @@ -2547,9 +2387,7 @@ async def test_get_team_alias_with_nested_fields(): assert jwt_handler.get_team_alias(nested_token, None) == "flat-team" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_alias_jwt_field="nonexistent.field") assert jwt_handler.get_team_alias(nested_token, "default-team") == "default-team" # Test with team_alias_jwt_field not configured @@ -2580,9 +2418,7 @@ async def test_is_required_team_id_with_team_alias_field(): assert jwt_handler.is_required_team_id() is True # Both fields set - should return True - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - team_id_jwt_field="team_id", team_alias_jwt_field="team_name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_id_jwt_field="team_id", team_alias_jwt_field="team_name") assert jwt_handler.is_required_team_id() is True @@ -2745,9 +2581,7 @@ async def test_get_org_alias_with_nested_fields(): } # Test nested access - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="company.organization.name" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="company.organization.name") assert jwt_handler.get_org_alias(nested_token, None) == "acme-corp" # Test flat access @@ -2755,9 +2589,7 @@ async def test_get_org_alias_with_nested_fields(): assert jwt_handler.get_org_alias(nested_token, None) == "flat-org" # Test missing field returns default - jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( - org_alias_jwt_field="nonexistent.field" - ) + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(org_alias_jwt_field="nonexistent.field") assert jwt_handler.get_org_alias(nested_token, "default-org") == "default-org" # Test with org_alias_jwt_field not configured @@ -2795,9 +2627,7 @@ async def test_get_objects_resolves_org_by_name(): models=[], ) - with patch( - "litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock - ) as mock_get_by_alias: + with patch("litellm.proxy.auth.handle_jwt.get_org_object_by_alias", new_callable=AsyncMock) as mock_get_by_alias: mock_get_by_alias.return_value = org_object ( @@ -2874,9 +2704,7 @@ async def test_resolve_jwks_url_resolves_oidc_discovery_document(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -2907,9 +2735,7 @@ async def test_resolve_jwks_url_caches_resolved_jwks_uri(): litellm_jwtauth=LiteLLM_JWTAuth(), ) - discovery_url = ( - "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" - ) + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" mock_response = MagicMock() @@ -3053,9 +2879,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_notation(): error_msg = str(exc_info.value) # Should mention the bad field name and suggest the fix assert "roles.0" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -3083,9 +2907,7 @@ async def test_find_and_validate_specific_team_id_hints_bracket_index_notation() error_msg = str(exc_info.value) assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}" - assert "roles" in error_msg and "list" in error_msg, ( - f"Expected hint about using 'roles' instead: {error_msg}" - ) + assert "roles" in error_msg and "list" in error_msg, f"Expected hint about using 'roles' instead: {error_msg}" @pytest.mark.asyncio @@ -3380,9 +3202,7 @@ async def test_auth_builder_single_team_fallback_membership_outage_raises_instea ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -3423,9 +3243,7 @@ def _reset_unscoped_warning_flag(): JWTHandler._unscoped_jwt_warning_emitted = False -def test_build_decode_kwargs_no_env_disables_both_verifications( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_no_env_disables_both_verifications(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3436,9 +3254,7 @@ def test_build_decode_kwargs_no_env_disables_both_verifications( assert kwargs["options"] == {"verify_aud": False, "verify_iss": False} -def test_build_decode_kwargs_audience_only_enables_aud_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_audience_only_enables_aud_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.delenv("JWT_ISSUER", raising=False) @@ -3450,9 +3266,7 @@ def test_build_decode_kwargs_audience_only_enables_aud_verification( assert kwargs["options"] == {"verify_iss": False} -def test_build_decode_kwargs_issuer_only_enables_iss_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_issuer_only_enables_iss_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.delenv("JWT_AUDIENCE", raising=False) monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3463,9 +3277,7 @@ def test_build_decode_kwargs_issuer_only_enables_iss_verification( assert kwargs["options"] == {"verify_aud": False} -def test_build_decode_kwargs_both_set_enables_full_verification( - monkeypatch, _reset_unscoped_warning_flag -): +def test_build_decode_kwargs_both_set_enables_full_verification(monkeypatch, _reset_unscoped_warning_flag): monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") @@ -3477,9 +3289,7 @@ def test_build_decode_kwargs_both_set_enables_full_verification( assert kwargs["options"] is None -def test_build_decode_kwargs_warns_once_when_unscoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_warns_once_when_unscoped(monkeypatch, _reset_unscoped_warning_flag, caplog): """The warning about unscoped JWT auth should fire on the first call but not on every subsequent decode.""" import logging @@ -3495,17 +3305,12 @@ def test_build_decode_kwargs_warns_once_when_unscoped( matching = [ r for r in caplog.records - if "JWT auth is enabled" in r.getMessage() - and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + if "JWT auth is enabled" in r.getMessage() and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] - assert len(matching) == 1, ( - f"Expected exactly one warning across 3 calls, got {len(matching)}" - ) + assert len(matching) == 1, f"Expected exactly one warning across 3 calls, got {len(matching)}" -def test_build_decode_kwargs_no_warning_when_scoped( - monkeypatch, _reset_unscoped_warning_flag, caplog -): +def test_build_decode_kwargs_no_warning_when_scoped(monkeypatch, _reset_unscoped_warning_flag, caplog): import logging monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") @@ -3514,11 +3319,7 @@ def test_build_decode_kwargs_no_warning_when_scoped( JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert matching == [] @@ -3576,11 +3377,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_returns_none( from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3652,9 +3449,7 @@ async def test_find_and_validate_specific_team_id_non_404_http_exception_propaga "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, ) as mock_get_team: - mock_get_team.side_effect = HTTPException( - status_code=status_code, detail="non-404 failure" - ) + mock_get_team.side_effect = HTTPException(status_code=status_code, detail="non-404 failure") with pytest.raises(HTTPException) as exc_info: await JWTAuthManager.find_and_validate_specific_team_id( @@ -3727,9 +3522,7 @@ async def test_find_team_with_model_access_resolved_team_without_model_still_rai async def mock_get_team_object(*_args, **_kwargs): return team - monkeypatch.setattr( - "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object - ) + monkeypatch.setattr("litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object) jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() @@ -3794,11 +3587,7 @@ async def test_find_team_with_model_access_unresolved_group_claim_default_raises from litellm.router import Router - router = Router( - model_list=[ - {"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}} - ] - ) + router = Router(model_list=[{"model_name": "gpt-4o-mini", "litellm_params": {"model": "gpt-4o-mini"}}]) proxy_server_module = types.ModuleType("proxy_server") proxy_server_module.llm_router = router monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) @@ -3835,12 +3624,7 @@ def test_canonical_user_id_rebinds_to_legacy_uuid(): jwt_email = "matt@example.com" user_object = LiteLLM_UserTable(user_id=legacy_uuid, user_email=jwt_email) - assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id=jwt_email, user_object=user_object - ) - == legacy_uuid - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=jwt_email, user_object=user_object) == legacy_uuid def test_canonical_user_id_no_change_when_ids_match(): @@ -3848,28 +3632,20 @@ def test_canonical_user_id_no_change_when_ids_match(): same = "alice@example.com" user_object = LiteLLM_UserTable(user_id=same, user_email=same) - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) - == same - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=same, user_object=user_object) == same def test_canonical_user_id_returns_claim_when_no_user_object(): """No resolved row (e.g. upsert disabled / brand new) -> keep the claim.""" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="newcomer@example.com", user_object=None - ) + JWTAuthManager._canonical_user_id_from_db(user_id="newcomer@example.com", user_object=None) == "newcomer@example.com" ) def test_canonical_user_id_returns_none_when_claim_none_and_no_object(): """Defensive: no claim and no row -> stays None, never invents an id.""" - assert ( - JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) - is None - ) + assert JWTAuthManager._canonical_user_id_from_db(user_id=None, user_object=None) is None def test_canonical_user_id_no_change_when_db_user_id_falsy(): @@ -3879,10 +3655,7 @@ def test_canonical_user_id_no_change_when_db_user_id_falsy(): user_id = "" assert ( - JWTAuthManager._canonical_user_id_from_db( - user_id="jwt@example.com", user_object=_Stub() - ) - == "jwt@example.com" + JWTAuthManager._canonical_user_id_from_db(user_id="jwt@example.com", user_object=_Stub()) == "jwt@example.com" ) @@ -3898,9 +3671,7 @@ async def test_auth_jwt_expired_token_raises_401_jwk_path(): jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3936,9 +3707,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): mock_cert.public_key.return_value.public_bytes.return_value = b"fake-key" with ( - patch.object( - jwt_handler, "get_public_key", new_callable=AsyncMock - ) as mock_get_public_key, + patch.object(jwt_handler, "get_public_key", new_callable=AsyncMock) as mock_get_public_key, patch( "litellm.proxy.auth.handle_jwt.jwt.get_unverified_header", return_value={"kid": "test-kid"}, @@ -3952,9 +3721,7 @@ async def test_auth_jwt_expired_token_raises_401_pem_cert_path(): side_effect=jwt_lib.ExpiredSignatureError("Signature has expired"), ), ): - mock_get_public_key.return_value = ( - "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" - ) + mock_get_public_key.return_value = "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----" with pytest.raises(ProxyException) as exc_info: await jwt_handler.auth_jwt(token="expired.jwt.token") @@ -4068,9 +3835,7 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): ) assert public_key == jwk - cached_keys = await cache.async_get_cache( - key="litellm_jwt_auth_keys_https://issuer.example.com/keys" - ) + cached_keys = await cache.async_get_cache(key="litellm_jwt_auth_keys_https://issuer.example.com/keys") assert cached_keys == [jwk] @@ -4315,9 +4080,7 @@ async def test_lowering_public_key_stale_ttl_stops_serving_a_copy_cached_under_t # The operator tightens the window and restarts; the cache, and its long-lived copy, survive. endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) - tightened = _get_jwt_handler_with_scripted_endpoint( - cache, endpoint, public_key_stale_ttl=lowered_stale_ttl - ) + tightened = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=lowered_stale_ttl) await cache.async_set_cache( key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", value=time.time() - 7200, @@ -4665,9 +4428,7 @@ def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) - assert ( - jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" - ) + assert jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" @pytest.mark.asyncio @@ -4786,9 +4547,7 @@ async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" - assert jwt_handler.get_team_id(token=claims, default_value=None) == ( - "example-org/litellm-fork" - ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == ("example-org/litellm-fork") @pytest.mark.asyncio @@ -4891,9 +4650,7 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert ( - jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" @pytest.mark.asyncio @@ -4926,7 +4683,7 @@ async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeyp kid="issuer-key", ) - with pytest.raises(Exception, match='Missing JWT Public Key URL from environment\\.') as exc: + with pytest.raises(Exception, match="Missing JWT Public Key URL from environment\\.") as exc: await jwt_handler.auth_jwt(token=token) assert "Missing JWT Public Key URL from environment." in str(exc.value) @@ -5000,7 +4757,7 @@ async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch) kid=shared_kid, ) - with pytest.raises(Exception, match='Validation fails: Signature verification failed') as exc: + with pytest.raises(Exception, match="Validation fails: Signature verification failed") as exc: await jwt_handler.auth_jwt(token=token) assert "Validation fails" in str(exc.value) @@ -5055,7 +4812,7 @@ def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='must configure audience or set') as exc: + with pytest.raises(Exception, match="must configure audience or set") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -5072,7 +4829,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): issuer = "https://issuer.example.com" jwks_url = f"{issuer}/keys" - with pytest.raises(Exception, match='cannot set audience and disable_audience_validation=True') as exc: + with pytest.raises(Exception, match="cannot set audience and disable_audience_validation=True") as exc: LiteLLM_JWTAuth( issuers=[ { @@ -5084,9 +4841,7 @@ def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): ] ) - assert "cannot set audience and disable_audience_validation=True together" in str( - exc.value - ) + assert "cannot set audience and disable_audience_validation=True together" in str(exc.value) @pytest.mark.asyncio @@ -5139,21 +4894,15 @@ async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): claims = await jwt_handler.auth_jwt(token=token) - assert jwt_handler.get_user_id(token=claims, default_value=None) == ( - "real-user@example.com" - ) - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_id(token=claims, default_value=None) == ("real-user@example.com") + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ "real-team", "secondary-team", ] assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" - assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( - "real-end-user" - ) + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ("real-end-user") @pytest.mark.asyncio @@ -5193,15 +4942,11 @@ async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims assert jwt_handler.get_user_id(token=claims, default_value=None) is None assert jwt_handler.get_team_id(token=claims, default_value=None) is None - assert jwt_handler.get_user_email(token=claims, default_value=None) == ( - "real-user@example.com" - ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ("real-user@example.com") @pytest.mark.asyncio -async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( - monkeypatch, caplog -): +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning(monkeypatch, caplog): import logging monkeypatch.delenv("JWT_AUDIENCE", raising=False) @@ -5251,11 +4996,7 @@ def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deploym JWTHandler._build_decode_kwargs() - matching = [ - r - for r in caplog.records - if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() - ] + matching = [r for r in caplog.records if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()] assert len(matching) == 1 @@ -5611,9 +5352,7 @@ async def test_resolve_db_team_fallback_skips_team_without_model_access(): teams=["restricted_team", "allowed_team"], ) teams = { - "restricted_team": LiteLLM_TeamTable( - team_id="restricted_team", models=["claude-3"] - ), + "restricted_team": LiteLLM_TeamTable(team_id="restricted_team", models=["claude-3"]), "allowed_team": LiteLLM_TeamTable(team_id="allowed_team", models=["gpt-4"]), } @@ -5796,9 +5535,7 @@ async def _run_auth_builder_with_header_team( jwt_handler = JWTHandler() jwt_handler.litellm_jwtauth = jwt_auth_config with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token - ), + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock, return_value=token), patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -5817,9 +5554,7 @@ async def _run_auth_builder_with_header_team( new_callable=AsyncMock, return_value=None, ), - patch.object( - JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids - ), + patch.object(JWTAuthManager, "get_all_team_ids", return_value=allowed_team_ids), patch.object( JWTAuthManager, "get_objects", @@ -5828,9 +5563,7 @@ async def _run_auth_builder_with_header_team( ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -5857,9 +5590,7 @@ async def _run_auth_builder_with_header_team( @pytest.mark.asyncio -async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> ( - None -): +async def test_auth_builder_header_team_not_found_matches_non_membership_denial() -> None: """A provisional x-litellm-team-id naming a nonexistent team must produce the exact same 403 shape as one naming an existing team outside the caller's memberships. Letting get_team_object's 404 surface would give any @@ -5879,19 +5610,15 @@ async def test_auth_builder_header_team_not_found_matches_non_membership_denial( return LiteLLM_TeamTable(team_id=team_id) with pytest.raises(HTTPException) as missing_exc: - await _run_auth_builder_with_header_team( - config, token, "team_ghost", user_object, _team_lookup_404, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_ghost", user_object, _team_lookup_404, set()) with pytest.raises(HTTPException) as outsider_exc: - await _run_auth_builder_with_header_team( - config, token, "team_other", user_object, team_exists, set() - ) + await _run_auth_builder_with_header_team(config, token, "team_other", user_object, team_exists, set()) assert missing_exc.value.status_code == 403 assert outsider_exc.value.status_code == 403 - assert missing_exc.value.detail.replace( - "team_ghost", "" - ) == outsider_exc.value.detail.replace("team_other", "") + assert missing_exc.value.detail.replace("team_ghost", "") == outsider_exc.value.detail.replace( + "team_other", "" + ) assert "exist" not in missing_exc.value.detail @@ -6088,9 +5815,7 @@ async def test_auth_builder_db_fallback_does_not_validate_rbac_team_against_db_m ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6242,9 +5967,7 @@ async def test_auth_builder_db_fallback_runs_when_only_team_id_default_set(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6326,9 +6049,7 @@ async def test_auth_builder_alias_only_token_resolves_alias_not_db_fallback(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6385,14 +6106,10 @@ async def test_find_and_validate_specific_team_id_alias_wins_over_team_id_defaul ) jwt_token = {"sub": "user-1", "team_alias": "my-team"} - alias_team = LiteLLM_TeamTable( - team_id="alias_resolved_team", team_alias="my-team" - ) + alias_team = LiteLLM_TeamTable(team_id="alias_resolved_team", team_alias="my-team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6439,9 +6156,7 @@ async def test_find_and_validate_specific_team_id_team_id_default_used_without_a default_team = LiteLLM_TeamTable(team_id="config_default_team") with ( - patch( - "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock - ) as mock_get_by_id, + patch("litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock) as mock_get_by_id, patch( "litellm.proxy.auth.handle_jwt.get_team_object_by_alias", new_callable=AsyncMock, @@ -6510,9 +6225,7 @@ async def test_auth_builder_db_fallback_enforces_passthrough_route_access(): ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6666,18 +6379,14 @@ async def test_sync_user_role_and_teams_singular_claim_reconciles_memberships(): "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { "team_stale_a", "team_stale_b", } - assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == { - "team_primary" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_add_user_to"]) == {"team_primary"} assert user.teams == ["team_primary"] @@ -6811,9 +6520,7 @@ async def test_auth_builder_header_cannot_override_rbac_team_under_db_fallback() ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -6863,9 +6570,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa async def call(route: str): with ( - patch.object( - jwt_handler, "auth_jwt", new_callable=AsyncMock - ) as mock_auth_jwt, + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), patch.object(jwt_handler, "get_rbac_role", return_value=None), patch.object(jwt_handler, "get_scopes", return_value=[]), @@ -6893,9 +6598,7 @@ async def test_auth_builder_header_team_enforces_team_allowed_routes_under_db_fa ), patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), patch.object(JWTAuthManager, "validate_object_id", return_value=True), - patch.object( - JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock - ), + patch.object(JWTAuthManager, "sync_user_role_and_teams", new_callable=AsyncMock), patch( "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock, @@ -7377,14 +7080,10 @@ async def test_sync_user_role_and_teams_singular_claim_only_recognized_under_fla "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", new_callable=AsyncMock, ) as mock_patch: - await JWTAuthManager.sync_user_role_and_teams( - jwt_handler, token, user, AsyncMock() - ) + await JWTAuthManager.sync_user_role_and_teams(jwt_handler, token, user, AsyncMock()) mock_patch.assert_awaited_once() - assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == { - "team_existing" - } + assert set(mock_patch.call_args.kwargs["teams_ids_to_remove_user_from"]) == {"team_existing"} assert mock_patch.call_args.kwargs["teams_ids_to_add_user_to"] == [] assert user.teams == [] @@ -7609,8 +7308,12 @@ async def test_auth_builder_propagates_agent_id_from_jwt_claim(monkeypatch, is_a if identity_only: identity = await JWTAuthManager.resolve_identity( - api_key=token, jwt_handler=jwt_handler, prisma_client=None, - user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + api_key=token, + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, ) assert identity.agent_id == "canonical-agent-id" return @@ -7645,8 +7348,12 @@ async def test_auth_builder_denies_jwt_naming_unregistered_agent_before_admin_ch if identity_only: with pytest.raises(HTTPException) as denial: await JWTAuthManager.resolve_identity( - api_key=token, jwt_handler=jwt_handler, prisma_client=None, - user_api_key_cache=None, parent_otel_span=None, proxy_logging_obj=None, + api_key=token, + jwt_handler=jwt_handler, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, ) assert denial.value.status_code == 403 return @@ -7713,7 +7420,9 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc from litellm.proxy.management_endpoints import team_endpoints handler, token = _entra_signed_app_token( - monkeypatch, azp="canonical-agent-id", scope=LiteLLM_JWTAuth().admin_jwt_scope, + monkeypatch, + azp="canonical-agent-id", + scope=LiteLLM_JWTAuth().admin_jwt_scope, ) handler.bind_agent_lookup(_entra_agent_registry()) handler.litellm_jwtauth.team_id_upsert = True @@ -7725,10 +7434,16 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc resolve = JWTAuthManager.auth_builder if admission else JWTAuthManager.authorize_jwt result = await resolve( - api_key=token, jwt_handler=handler, request_data={}, general_settings={}, - route="/chat/completions", prisma_client=database, - user_api_key_cache=handler.user_api_key_cache, parent_otel_span=None, - proxy_logging_obj=MagicMock(), request_headers={"x-litellm-team-id": "new-team"}, + api_key=token, + jwt_handler=handler, + request_data={}, + general_settings={}, + route="/chat/completions", + prisma_client=database, + user_api_key_cache=handler.user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + request_headers={"x-litellm-team-id": "new-team"}, ) assert result["is_proxy_admin"] is True @@ -7742,16 +7457,20 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc def _explicit_identity_registry() -> AgentRegistry: registry: Final = AgentRegistry() - registry.register_agent(AgentResponse( - agent_id="explicit-agent-id", - agent_name="Readable agent name", - agent_card_params={}, - litellm_params={"identity": { - "provider": "microsoft_entra", - "tenant_id": "11111111-1111-4111-8111-111111111111", - "client_id": "22222222-2222-4222-8222-222222222222", - }}, - )) + registry.register_agent( + AgentResponse( + agent_id="explicit-agent-id", + agent_name="Readable agent name", + agent_card_params={}, + litellm_params={ + "identity": { + "provider": "microsoft_entra", + "tenant_id": "11111111-1111-4111-8111-111111111111", + "client_id": "22222222-2222-4222-8222-222222222222", + } + }, + ) + ) return registry @@ -7772,23 +7491,30 @@ def test_runtime_json_cannot_establish_a_managed_identity(claim_field: str | Non assert failure.value.status_code == 403 -@pytest.mark.parametrize("override", [ - {"iss": "https://attacker.example"}, - {"tid": "33333333-3333-4333-8333-333333333333"}, - {"azp": "33333333-3333-4333-8333-333333333333"}, - {"azp": "explicit-agent-id"}, - {"azp": "Readable agent name"}, -]) +@pytest.mark.parametrize( + "override", + [ + {"iss": "https://attacker.example"}, + {"tid": "33333333-3333-4333-8333-333333333333"}, + {"azp": "33333333-3333-4333-8333-333333333333"}, + {"azp": "explicit-agent-id"}, + {"azp": "Readable agent name"}, + ], +) def test_explicit_entra_identity_cannot_be_claimed_via_legacy_lookup(override: Mapping[str, object]) -> None: registry: Final = _explicit_identity_registry() handler: Final = _entra_agent_jwt_handler("azp") with pytest.raises(HTTPException) as failure: - JWTAuthManager.resolve_agent_id(handler, { - "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0", - "tid": "11111111-1111-4111-8111-111111111111", - "azp": "22222222-2222-4222-8222-222222222222", - **override, - }, registry) + JWTAuthManager.resolve_agent_id( + handler, + { + "iss": "https://login.microsoftonline.com/11111111-1111-4111-8111-111111111111/v2.0", + "tid": "11111111-1111-4111-8111-111111111111", + "azp": "22222222-2222-4222-8222-222222222222", + **override, + }, + registry, + ) assert failure.value.status_code == 403 @@ -7805,7 +7531,12 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning cache: Final = UserApiKeyCache() cache.set_cache("litellm_jwt_auth_keys_https://admin.example/jwks", [jwk]) user_id: Final = f"admin-status-{existing_user}-{warm_cache}-{email}" - user: Final = LiteLLM_UserTable(user_id=user_id, user_email="admin@allowed.example", metadata={"scim_active": False}, organization_memberships=[]) + user: Final = LiteLLM_UserTable( + user_id=user_id, + user_email="admin@allowed.example", + metadata={"scim_active": False}, + organization_memberships=[], + ) if existing_user and warm_cache: cache.set_cache(user_id, user) database: Final = MagicMock() @@ -7818,7 +7549,9 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning prisma_client=database, user_api_key_cache=cache, litellm_jwtauth=LiteLLM_JWTAuth( - user_id_jwt_field="sub", user_id_upsert=True, user_email_jwt_field="email", + user_id_jwt_field="sub", + user_id_upsert=True, + user_email_jwt_field="email", user_allowed_email_domain="allowed.example", ), ) @@ -7826,12 +7559,22 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning monkeypatch.setenv("JWT_ISSUER", "https://admin.example") monkeypatch.setenv("JWT_AUDIENCE", "gateway") token: Final = _encode_rsa_jwt( - private_key, "https://admin.example", "gateway", "admin-status", + private_key, + "https://admin.example", + "gateway", + "admin-status", {"sub": user_id, "scope": "litellm_proxy_admin", **({"email": email} if email else {})}, ) result: Final = await JWTAuthManager.auth_builder( - api_key=token, jwt_handler=handler, prisma_client=database, user_api_key_cache=cache, - parent_otel_span=None, proxy_logging_obj=MagicMock(), request_data={}, general_settings={}, route="/user/info", + api_key=token, + jwt_handler=handler, + prisma_client=database, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + request_data={}, + general_settings={}, + route="/user/info", ) assert result["is_proxy_admin"] is True assert result["user_id"] == user_id @@ -7841,17 +7584,183 @@ async def test_scope_admin_admission_resolves_existing_user_without_provisioning assert users.find_unique.await_count == (0 if warm_cache else 1) +# --- B5: JWT algorithm allowlists (LIT-8429) --- + + +def _okp_keypair_and_jwk() -> "tuple[object, dict]": + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + private_key = ed25519.Ed25519PrivateKey.generate() + jwk = { + **json.loads(__import__("jwt").algorithms.OKPAlgorithm.to_jwk(private_key.public_key())), + "kid": "ed", + "alg": "EdDSA", + "use": "sig", + } + return private_key, jwk + + +def _eddsa_jwt(private_key: object, kid: str = "ed") -> str: + import jwt + + current_time = int(time.time()) + return jwt.encode( + {"sub": "test-subject", "iat": current_time, "exp": current_time + 300}, + private_key, # pyright: ignore[reportArgumentType] + algorithm="EdDSA", + headers={"kid": kid}, + ) + + +def test_approved_jwt_algorithm_literal_matches_tuple(): + from typing import get_args + + from litellm.types.proxy.auth.jwt_algorithms import APPROVED_JWT_ALGORITHMS, ApprovedJwtAlgorithm + + assert set(get_args(ApprovedJwtAlgorithm)) == set(APPROVED_JWT_ALGORITHMS) + + +def test_allowed_jwt_algorithms_drops_legacy_only_in_fips_mode(): + from litellm.proxy.auth.jwt_algorithms import allowed_jwt_algorithms + + assert "EdDSA" not in allowed_jwt_algorithms(True) + assert list(allowed_jwt_algorithms(False)) == JWTHandler.SUPPORTED_JWT_ALGORITHMS + + +@pytest.mark.parametrize( + "key,algorithms,expected", + [ + ({"kty": "RSA", "alg": "RS256", "kid": "a"}, ("RS256",), True), + ({"kty": "RSA", "alg": "HS256", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "OKP", "alg": "EdDSA", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "RSA", "kid": "a"}, ("RS256", "ES256"), True), + ({"kty": "EC", "kid": "a"}, ("RS256", "ES256"), True), + ({"kty": "EC", "kid": "a"}, ("RS256",), False), + ({"kty": "OKP", "kid": "a"}, ("RS256", "EdDSA"), True), + ({"kty": "OKP", "kid": "a"}, ("RS256", "ES256"), False), + ({"kty": "oct", "kid": "a"}, ("RS256", "HS256"), False), + ({"kid": "a"}, ("RS256",), False), + ], +) +def test_jwks_keys_for_filters_by_declared_or_inferred_algorithm(key, algorithms, expected): + from litellm.proxy.auth.jwt_algorithms import jwks_keys_for + + assert list(jwks_keys_for([key], algorithms)) == ([key] if expected else []) + + +def test_fips_mode_rejects_eddsa_token_but_accepts_rs256(): + import jwt + + ed_private, ed_jwk = _okp_keypair_and_jwk() + rsa_private, rsa_jwk = _rsa_keypair_and_jwk() + handler = JWTHandler(fips_mode=lambda: True) + + with pytest.raises(jwt.exceptions.InvalidAlgorithmError): + handler._decode_jwt_with_public_key( + token=_eddsa_jwt(ed_private), + public_key=ed_jwk, + audience=None, + disable_audience_validation=True, + ) + claims: Final = handler._decode_jwt_with_public_key( + token=_encode_rsa_jwt(rsa_private, "iss", "aud", "rsa"), + public_key=rsa_jwk, + audience="aud", + issuer="iss", + ) + assert claims["sub"] == "test-subject" + + +def test_eddsa_token_accepted_with_deprecation_log_outside_fips(caplog): + import logging + + from litellm.proxy.auth.handle_jwt import _log_eddsa_deprecation + + ed_private, ed_jwk = _okp_keypair_and_jwk() + handler = JWTHandler(fips_mode=lambda: False) + + _log_eddsa_deprecation.cache_clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + claims: Final = handler._decode_jwt_with_public_key( + token=_eddsa_jwt(ed_private), + public_key=ed_jwk, + audience=None, + disable_audience_validation=True, + ) + assert claims["sub"] == "test-subject" + assert "EdDSA" in caplog.text and "deprecated" in caplog.text and "LITELLM_FIPS_MODE" in caplog.text + + _log_eddsa_deprecation.cache_clear() + caplog.clear() + rsa_private, rsa_jwk = _rsa_keypair_and_jwk() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + handler._decode_jwt_with_public_key( + token=_encode_rsa_jwt(rsa_private, "iss", "aud", "rsa"), + public_key=rsa_jwk, + audience="aud", + issuer="iss", + ) + assert "deprecated" not in caplog.text + _log_eddsa_deprecation.cache_clear() + + +def _rsa_keypair_and_jwk() -> "tuple[object, dict]": + import json + + import jwt as jwt_lib + from cryptography.hazmat.primitives.asymmetric import rsa + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + jwk = { + **json.loads(jwt_lib.algorithms.RSAAlgorithm.to_jwk(private_key.public_key())), + "kid": "rsa", + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +@pytest.mark.asyncio +async def test_jwks_url_kid_matching_only_a_filtered_key_raises(): + _, ed_jwk = _okp_keypair_and_jwk() + jwks_url = "https://idp.example.com/jwks" + + fips_handler = JWTHandler(fips_mode=lambda: True) + fips_handler.user_api_key_cache = DualCache() + await fips_handler.user_api_key_cache.async_set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", value=[ed_jwk], ttl=600 + ) + with pytest.raises(NoMatchingJWTPublicKeyError): + await fips_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="ed") + + normal_handler = JWTHandler(fips_mode=lambda: False) + normal_handler.user_api_key_cache = DualCache() + await normal_handler.user_api_key_cache.async_set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", value=[ed_jwk], ttl=600 + ) + assert (await normal_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="ed"))["kid"] == "ed" + + @pytest.mark.asyncio @pytest.mark.parametrize("mode", ["autonomous", "both", "delegated"]) @pytest.mark.parametrize("audience_validation", (True, False)) @pytest.mark.parametrize( "route,allowed", [ - ("/chat/completions", True), ("/v1/messages", True), ("/v1/responses", True), - ("/mcp-rest/tools/call", True), ("/a2a/target", True), - ("/v1/files", False), ("/v1/batches", False), ("/v1/vector_stores", False), - ("/v1/containers", False), ("/openai/v1/files", False), - ("/v1/responses/other-response", False), ("/v1/realtime/client_secrets", False), + ("/chat/completions", True), + ("/v1/messages", True), + ("/v1/responses", True), + ("/mcp-rest/tools/call", True), + ("/a2a/target", True), + ("/v1/files", False), + ("/v1/batches", False), + ("/v1/vector_stores", False), + ("/v1/containers", False), + ("/openai/v1/files", False), + ("/v1/responses/other-response", False), + ("/v1/realtime/client_secrets", False), ], ) async def test_managed_application_uses_persisted_identity_without_provisioning_human( @@ -7990,7 +7899,9 @@ async def test_database_free_jwt_admission_with_entra_shaped_claims(monkeypatch: ) ) handler: Final = JWTHandler() - handler.update_environment(None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"])) + handler.update_environment( + None, cache, LiteLLM_JWTAuth(agent_id_jwt_field="agent", admin_allowed_routes=["llm_api_routes"]) + ) handler.bind_agent_lookup(registry) token: Final = _encode_rsa_jwt( private_key, @@ -8073,11 +7984,20 @@ async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy( monkeypatch.setenv("JWT_ISSUER", issuer) monkeypatch.setenv("JWT_AUDIENCE", "gateway") binding: Final = AgentIdentityBinding( - agent_id="managed", provider="microsoft_entra", issuer=issuer, tenant_id="tenant", - client_id="client", service_principal_id="principal", revision="current", + agent_id="managed", + provider="microsoft_entra", + issuer=issuer, + tenant_id="tenant", + client_id="client", + service_principal_id="principal", + revision="current", ) agent: Final = AgentResponse( - agent_id="managed", agent_name="Managed", agent_card_params={}, identity_managed=True, identity=binding, + agent_id="managed", + agent_name="Managed", + agent_card_params={}, + identity_managed=True, + identity=binding, ) database: Final = MagicMock() database.writer_db.litellm_agentidentity.find_unique = AsyncMock(return_value=binding) @@ -8089,8 +8009,15 @@ async def test_managed_jwt_reuses_binding_lookup_but_rechecks_disabled_policy( private_key, issuer, "gateway", "managed-cache", {"tid": "tenant", "azp": "client", "oid": "principal"} ) arguments: Final = dict( - api_key=token, jwt_handler=handler, request_data={}, general_settings={}, route="/chat/completions", - prisma_client=database, user_api_key_cache=cache, parent_otel_span=None, proxy_logging_obj=MagicMock(), + api_key=token, + jwt_handler=handler, + request_data={}, + general_settings={}, + route="/chat/completions", + prisma_client=database, + user_api_key_cache=cache, + parent_otel_span=None, + proxy_logging_obj=MagicMock(), ) for _ in range(2): result: Final = await JWTAuthManager.authorize_jwt(**arguments) diff --git a/tests/unit/proxy/guardrails/test_mcp_jwt_signer.py b/tests/unit/proxy/guardrails/test_mcp_jwt_signer.py index a7b24169398..7d904b87cd9 100644 --- a/tests/unit/proxy/guardrails/test_mcp_jwt_signer.py +++ b/tests/unit/proxy/guardrails/test_mcp_jwt_signer.py @@ -1292,3 +1292,101 @@ async def test_inject_mcp_jwt_signs_for_tool_call_path(): scopes = set(decoded["scope"].split()) assert "mcp:tools/call" in scopes assert "mcp:tools/search_web:call" in scopes + + +# --------------------------------------------------------------------------- +# B5: incoming-JWT JWKS allowlist (LIT-8429) +# --------------------------------------------------------------------------- + + +def _okp_idp_key_and_token(now: int): + import json + + from cryptography.hazmat.primitives.asymmetric import ed25519 + + private_key = ed25519.Ed25519PrivateKey.generate() + jwk = { + **json.loads(jwt.algorithms.OKPAlgorithm.to_jwk(private_key.public_key())), + "kid": "ed", + "alg": "EdDSA", + } + token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + private_key, + algorithm="EdDSA", + headers={"kid": "ed"}, + ) + return jwk, token + + +def _oct_idp_key_and_token(now: int): + secret = b"integration-hs256-client-secret-0123456789abcdef" + jwk = { + "kty": "oct", + "kid": "sym", + "alg": "HS256", + "k": base64.urlsafe_b64encode(secret).rstrip(b"=").decode(), + } + token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + secret, + algorithm="HS256", + headers={"kid": "sym"}, + ) + return jwk, token + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "key_and_token", (_oct_idp_key_and_token, _okp_idp_key_and_token), ids=("oct-HS256", "OKP-EdDSA") +) +async def test_verify_incoming_jwt_rejects_jwks_without_approved_algorithms(key_and_token): + """A JWKS whose only keys use non-approved algorithms can never verify the incoming token.""" + jwks_key, incoming_token = key_and_token(int(time.time())) + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration", + ) + with patch.object( + signer, + "_get_oidc_discovery", + new_callable=AsyncMock, + return_value={"jwks_uri": "https://idp.example.com/jwks"}, + ): + with patch( + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks", + new_callable=AsyncMock, + return_value=[jwks_key], + ): + with pytest.raises(jwt.exceptions.PyJWKSetError, match="approved signing algorithm"): + await signer._verify_incoming_jwt(incoming_token) + + +@pytest.mark.asyncio +async def test_verify_incoming_jwt_ignores_filtered_keys_sharing_kid(): + """An approved RS256 key still verifies when the JWKS also carries a non-approved key with the same kid.""" + signer = _make_signer( + access_token_discovery_uri="https://idp.example.com/.well-known/openid-configuration", + ) + now = int(time.time()) + oct_jwk, _ = _oct_idp_key_and_token(now) + oct_jwk["kid"] = signer._kid + incoming_token = jwt.encode( + {"sub": "idp-user", "iat": now, "exp": now + 300}, + signer._private_key, + algorithm="RS256", + headers={"kid": signer._kid}, + ) + jwks = [oct_jwk, *signer.get_jwks()["keys"]] + with patch.object( + signer, + "_get_oidc_discovery", + new_callable=AsyncMock, + return_value={"jwks_uri": "https://idp.example.com/jwks"}, + ): + with patch( + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer._fetch_jwks", + new_callable=AsyncMock, + return_value=jwks, + ): + payload = await signer._verify_incoming_jwt(incoming_token) + assert payload["sub"] == "idp-user" diff --git a/tests/unit/proxy/test__types.py b/tests/unit/proxy/test__types.py index 70c5a153647..c723631a16c 100644 --- a/tests/unit/proxy/test__types.py +++ b/tests/unit/proxy/test__types.py @@ -387,6 +387,38 @@ def test_change_password_request_passwords_hidden_from_repr(): for rendered in (repr(request), str(request)): assert "hunter2hunter2" not in rendered assert "NewP@ssw0rd-2026" not in rendered + + +def test_mcp_server_requests_reject_non_approved_client_assertion_signing_alg() -> None: + """The REST boundary is strict even though the stored blob stays lenient: + a write carrying HS256/hs256/"" must 422 naming the field.""" + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + for cls, base in ( + (NewMCPServerRequest, {"transport": "http", "url": "https://mcp.example.com"}), + (UpdateMCPServerRequest, {"server_id": "srv-1", "transport": "http", "url": "https://mcp.example.com"}), + ): + for alg in ("HS256", "hs256", "", "EdDSA"): + with pytest.raises(ValidationError) as exc: + cls(**base, credentials={"client_assertion_signing_alg": alg}) + assert "client_assertion_signing_alg" in str(exc.value), f"{cls.__name__} accepted {alg!r}" + + +def test_mcp_server_requests_accept_approved_client_assertion_signing_alg() -> None: + from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest + + for alg in ("ES256", "PS384", "RS256", None): + request = NewMCPServerRequest( + transport="http", + url="https://mcp.example.com", + credentials={"client_assertion_signing_alg": alg}, + ) + assert request.credentials is not None + assert request.credentials["client_assertion_signing_alg"] == alg + update = UpdateMCPServerRequest(server_id="srv-1") + assert update.credentials is None + + @pytest.mark.parametrize("versions", [[], ["2099-01-01"], ["2026-07-28"]]) def test_mcp_advertised_versions_reject_unavailable_revisions(versions): from pydantic import ValidationError @@ -401,7 +433,12 @@ def test_mcp_advertised_versions_reject_unavailable_revisions(versions): def test_mcp_metadata_rejects_unavailable_upstream_protocol(revision): from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest - payload = {"server_id": "test", "transport": "http", "url": "https://example.com/mcp", "mcp_info": {"protocol_version": revision}} + payload = { + "server_id": "test", + "transport": "http", + "url": "https://example.com/mcp", + "mcp_info": {"protocol_version": revision}, + } for model in (NewMCPServerRequest, UpdateMCPServerRequest): with pytest.raises(ValidationError): model.model_validate(payload) @@ -457,7 +494,10 @@ def test_an_enabled_stdio_mcp_server_still_needs_a_command_and_args(monkeypatch, def test_an_http_mcp_server_is_unaffected_by_the_stdio_flag(monkeypatch, request_model): monkeypatch.delenv("LITELLM_ENABLE_MCP_STDIO", raising=False) - assert request_model(server_id="http-1", transport="http", url="https://mcp.example.com").url == "https://mcp.example.com" + assert ( + request_model(server_id="http-1", transport="http", url="https://mcp.example.com").url + == "https://mcp.example.com" + ) with pytest.raises(ValidationError, match="url or spec_path is required"): request_model(server_id="http-1", transport="http") @@ -470,9 +510,13 @@ def test_a_non_mapping_mcp_server_payload_gets_a_validation_error(request_model) @pytest.mark.parametrize("request_model", MCP_SERVER_REQUESTS) def test_modern_http_upstream_protocol_is_available(request_model): - parsed = request_model.model_validate({ - "server_id": "modern", "transport": "http", "url": "https://example.com/mcp", - "mcp_info": {"protocol_version": "2026-07-28"}, - }) + parsed = request_model.model_validate( + { + "server_id": "modern", + "transport": "http", + "url": "https://example.com/mcp", + "mcp_info": {"protocol_version": "2026-07-28"}, + } + ) assert parsed.mcp_info["protocol_version"] == "2026-07-28" assert parsed.transport == "http"