mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-27 01:22:18 +00:00
feat(proxy): support issuer-scoped JWT auth
This commit is contained in:
parent
c279f4e3ef
commit
1d731af5df
3 changed files with 247 additions and 18 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue