From 1d731af5dfc24c15098f8a83f02d811b5d3eccb3 Mon Sep 17 00:00:00 2001 From: gym-cmd <186399764+gym-cmd@users.noreply.github.com> Date: Fri, 15 May 2026 15:04:31 +0100 Subject: [PATCH] feat(proxy): support issuer-scoped JWT auth --- litellm/proxy/_types.py | 6 +- litellm/proxy/auth/handle_jwt.py | 81 ++++++++++--- tests/proxy_unit_tests/test_jwt.py | 178 +++++++++++++++++++++++++++++ 3 files changed, 247 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e49490df4d0..f8387ec07e9 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4348,7 +4348,11 @@ class JWTIssuerConfig(BaseModel): ) audience: Optional[Union[str, List[str]]] = Field( default=None, - description="Expected token audience for this issuer. If omitted, audience validation is disabled for this issuer.", + description="Expected token audience for this issuer.", + ) + disable_audience_validation: bool = Field( + default=False, + description="Explicitly disable audience validation for this issuer. Use only when the issuer cannot provide an audience suitable for LiteLLM.", ) user_id_jwt_field: Optional[str] = Field( default=None, diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 36ccf58a40d..f47a307cf72 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -99,6 +99,15 @@ class JWTHandler: LITELLM_TEAM_IDS_CLAIM = "_litellm_team_ids" LITELLM_ORG_ID_CLAIM = "_litellm_org_id" LITELLM_END_USER_ID_CLAIM = "_litellm_end_user_id" + LITELLM_INTERNAL_CLAIMS = ( + LITELLM_JWT_ISSUER_CLAIM, + LITELLM_USER_ID_CLAIM, + LITELLM_USER_EMAIL_CLAIM, + LITELLM_TEAM_ID_CLAIM, + LITELLM_TEAM_IDS_CLAIM, + LITELLM_ORG_ID_CLAIM, + LITELLM_END_USER_ID_CLAIM, + ) def __init__( self, @@ -221,12 +230,27 @@ class JWTHandler: return True return False + def _is_trusted_issuer_normalized_token(self, token: dict) -> bool: + issuer = token.get(self.LITELLM_JWT_ISSUER_CLAIM) + if not isinstance(issuer, str) or not issuer: + return False + + litellm_jwtauth = getattr(self, "litellm_jwtauth", None) + issuer_configs = getattr(litellm_jwtauth, "issuers", None) or [] + return any(issuer_config.issuer == issuer for issuer_config in issuer_configs) + + def _has_trusted_issuer_normalized_claim(self, token: dict, claim: str) -> bool: + return self._is_trusted_issuer_normalized_token(token=token) and claim in token + def get_team_ids_from_jwt(self, token: dict) -> List[str]: - issuer_team_ids = token.get(self.LITELLM_TEAM_IDS_CLAIM) - if isinstance(issuer_team_ids, list): - return issuer_team_ids - if isinstance(issuer_team_ids, str): - return [issuer_team_ids] + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_TEAM_IDS_CLAIM + ): + issuer_team_ids = token.get(self.LITELLM_TEAM_IDS_CLAIM) + if isinstance(issuer_team_ids, list): + return issuer_team_ids + if isinstance(issuer_team_ids, str): + return [issuer_team_ids] if self.litellm_jwtauth.team_ids_jwt_field is not None: team_ids: Optional[List[str]] = get_nested_value( @@ -241,7 +265,9 @@ class JWTHandler: def get_end_user_id( self, token: dict, default_value: Optional[str] ) -> Optional[str]: - if self.LITELLM_END_USER_ID_CLAIM in token: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_END_USER_ID_CLAIM + ): return token.get(self.LITELLM_END_USER_ID_CLAIM) try: @@ -285,7 +311,9 @@ class JWTHandler: return False def get_team_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: - if self.LITELLM_TEAM_ID_CLAIM in token: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_TEAM_ID_CLAIM + ): team_id = token.get(self.LITELLM_TEAM_ID_CLAIM) if isinstance(team_id, list): return team_id[0] if team_id else default_value @@ -364,7 +392,9 @@ class JWTHandler: return self.litellm_jwtauth.user_id_upsert def get_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: - if self.LITELLM_USER_ID_CLAIM in token: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_USER_ID_CLAIM + ): return token.get(self.LITELLM_USER_ID_CLAIM) try: @@ -458,7 +488,9 @@ class JWTHandler: def get_user_email( self, token: dict, default_value: Optional[str] ) -> Optional[str]: - if self.LITELLM_USER_EMAIL_CLAIM in token: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_USER_EMAIL_CLAIM + ): return token.get(self.LITELLM_USER_EMAIL_CLAIM) try: @@ -489,7 +521,9 @@ class JWTHandler: return object_id def get_org_id(self, token: dict, default_value: Optional[str]) -> Optional[str]: - if self.LITELLM_ORG_ID_CLAIM in token: + if self._has_trusted_issuer_normalized_claim( + token=token, claim=self.LITELLM_ORG_ID_CLAIM + ): return token.get(self.LITELLM_ORG_ID_CLAIM) try: @@ -851,6 +885,10 @@ class JWTHandler: def _apply_issuer_claim_mappings( self, token: dict, issuer_config: JWTIssuerConfig ) -> dict: + source_token = {**token} + for claim in self.LITELLM_INTERNAL_CLAIMS: + token.pop(claim, None) + token[self.LITELLM_JWT_ISSUER_CLAIM] = issuer_config.issuer claim_mappings = [ (issuer_config.user_id_jwt_field, self.LITELLM_USER_ID_CLAIM), @@ -865,7 +903,7 @@ class JWTHandler: if source_claim is None: continue token[normalized_claim] = self._get_claim_value_for_issuer_mapping( - token=token, + token=source_token, claim_field=source_claim, issuer=issuer_config.issuer, ) @@ -932,6 +970,14 @@ class JWTHandler: async def _auth_jwt_with_issuer( self, token: str, issuer_config: JWTIssuerConfig, kid: Optional[str] ) -> dict: + if ( + issuer_config.audience is None + and not issuer_config.disable_audience_validation + ): + raise Exception( + f"JWT issuer {issuer_config.issuer} must configure audience or set disable_audience_validation=True" + ) + public_key = await self._get_public_key_from_jwks_url( jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config), kid=kid, @@ -943,18 +989,17 @@ class JWTHandler: audience=issuer_config.audience, issuer=issuer_config.issuer, ) - return self._apply_issuer_claim_mappings( - token=payload, - issuer_config=issuer_config, - ) except jwt.ExpiredSignatureError: raise Exception("Token Expired") except Exception as e: raise Exception(f"Validation fails: {str(e)}") - async def auth_jwt(self, token: str) -> dict: - decode_kwargs = self._build_decode_kwargs() + return self._apply_issuer_claim_mappings( + token=payload, + issuer_config=issuer_config, + ) + async def auth_jwt(self, token: str) -> dict: header = jwt.get_unverified_header(token) verbose_proxy_logger.debug("header: %s", header) @@ -969,6 +1014,8 @@ class JWTHandler: kid=kid, ) + decode_kwargs = self._build_decode_kwargs() + public_key = await self.get_public_key(kid=kid) if public_key is not None: diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index 633d6c579bd..8d375608f53 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -3,6 +3,7 @@ import asyncio import base64 +import logging import os import random import sys @@ -1736,6 +1737,7 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): "issuer": issuer, "jwks_url": jwks_url, "audience": None, + "disable_audience_validation": True, "user_id_jwt_field": "kubernetes\\.io.namespace", } ], @@ -1888,3 +1890,179 @@ async def test_multi_issuer_jwt_missing_mapped_claim_fails_closed(monkeypatch): await jwt_handler.auth_jwt(token=token) assert "missing required mapped claim: email" in str(exc.value) + assert "Validation fails" not in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="some-other-client", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "must configure audience" in str(exc.value) + + +@pytest.mark.asyncio +async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + jwks_url = "https://global-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + private_key, jwk = _get_rsa_key_and_jwk(kid="global-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="email", + user_email_jwt_field="email", + team_id_jwt_field="team.id", + team_ids_jwt_field="teams", + org_id_jwt_field="org.id", + end_user_id_jwt_field="end_user.id", + ), + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://global-issuer.example.com", + audience="some-other-client", + kid="global-key", + extra_claims={ + "email": "real-user@example.com", + "team": {"id": "real-team"}, + "teams": ["real-team", "secondary-team"], + "org": {"id": "real-org"}, + "end_user": {"id": "real-end-user"}, + JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://issuer.example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_USER_EMAIL_CLAIM: "victim@example.com", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + JWTHandler.LITELLM_TEAM_IDS_CLAIM: ["victim-team"], + JWTHandler.LITELLM_ORG_ID_CLAIM: "victim-org", + JWTHandler.LITELLM_END_USER_ID_CLAIM: "victim-end-user", + }, + ) + + 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_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" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_email_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + extra_claims={ + "email": "real-user@example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + 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" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + JWTHandler._unscoped_jwt_warning_emitted = False + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + with caplog.at_level(logging.WARNING): + await jwt_handler.auth_jwt(token=token) + + assert "Tokens minted by any application" not in caplog.text + assert JWTHandler._unscoped_jwt_warning_emitted is False