feat(proxy): support issuer-scoped JWT auth

This commit is contained in:
gym-cmd 2026-05-15 15:04:31 +01:00
parent c279f4e3ef
commit 1d731af5df
3 changed files with 247 additions and 18 deletions

View file

@ -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,

View file

@ -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:

View file

@ -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