feat(proxy): support issuer-scoped JWT auth

This commit is contained in:
gym-cmd 2026-05-15 13:41:28 +01:00
parent 8e864f33ee
commit da24c8983a
3 changed files with 614 additions and 92 deletions

View file

@ -2438,7 +2438,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
)
mcp_trusted_proxy_ranges: Optional[List[str]] = Field(
None,
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.",
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.",
)
trusted_proxy_ranges: Optional[List[str]] = Field(
None,
@ -4333,6 +4333,53 @@ class JWTRoutingOverride(BaseModel):
}
class JWTIssuerConfig(BaseModel):
"""
Issuer-bound JWT validation configuration.
When configured, LiteLLM selects this issuer by the token's unverified `iss`
claim, then validates the token only against this issuer's JWKS and audience.
"""
issuer: str = Field(description="Exact expected JWT issuer (`iss`) value.")
jwks_url: Optional[str] = Field(
default=None,
description="Issuer JWKS URL. If omitted, LiteLLM uses the issuer's OIDC discovery document.",
)
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.",
)
user_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's user id.",
)
user_email_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's user email.",
)
team_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's team id.",
)
team_ids_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's team ids.",
)
org_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's organization id.",
)
end_user_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's end-user id.",
)
model_config = {
"extra": "forbid",
}
class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"""
A class to define the roles and permissions for a LiteLLM Proxy w/ JWT Auth.
@ -4437,6 +4484,10 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
default=None,
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
)
issuers: Optional[List[JWTIssuerConfig]] = Field(
default=None,
description="Optional issuer-bound JWT validation rules. When set, tokens must match one configured issuer by exact `iss` claim before JWKS lookup.",
)
#########################################################
def __init__(self, **kwargs: Any) -> None:

View file

@ -12,7 +12,7 @@ import fnmatch
import hashlib
import os
import re
from typing import Any, List, Literal, Optional, Set, Tuple, cast
from typing import Any, List, Literal, Optional, Set, Tuple, Union, cast
from cryptography import x509
from cryptography.hazmat.backends import default_backend
@ -29,6 +29,7 @@ from litellm.proxy._types import (
RBAC_ROLES,
JWKKeyValue,
JWTAuthBuilderResult,
JWTIssuerConfig,
JWTKeyItem,
LiteLLM_EndUserTable,
LiteLLM_JWTAuth,
@ -91,6 +92,13 @@ class JWTHandler:
"ES512",
"EdDSA",
]
LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer"
LITELLM_USER_ID_CLAIM = "_litellm_user_id"
LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email"
LITELLM_TEAM_ID_CLAIM = "_litellm_team_id"
LITELLM_TEAM_IDS_CLAIM = "_litellm_team_ids"
LITELLM_ORG_ID_CLAIM = "_litellm_org_id"
LITELLM_END_USER_ID_CLAIM = "_litellm_end_user_id"
def __init__(
self,
@ -214,6 +222,12 @@ class JWTHandler:
return False
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.litellm_jwtauth.team_ids_jwt_field is not None:
team_ids: Optional[List[str]] = get_nested_value(
data=token,
@ -227,6 +241,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:
return token.get(self.LITELLM_END_USER_ID_CLAIM)
try:
if self.litellm_jwtauth.end_user_id_jwt_field is not None:
user_id = get_nested_value(
@ -268,6 +285,12 @@ 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:
team_id = token.get(self.LITELLM_TEAM_ID_CLAIM)
if isinstance(team_id, list):
return team_id[0] if team_id else default_value
return team_id
try:
if self.litellm_jwtauth.team_id_jwt_field is not None:
# Use a sentinel value to detect if the path actually exists
@ -341,6 +364,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:
return token.get(self.LITELLM_USER_ID_CLAIM)
try:
if self.litellm_jwtauth.user_id_jwt_field is not None:
user_id = get_nested_value(
@ -432,6 +458,9 @@ class JWTHandler:
def get_user_email(
self, token: dict, default_value: Optional[str]
) -> Optional[str]:
if self.LITELLM_USER_EMAIL_CLAIM in token:
return token.get(self.LITELLM_USER_EMAIL_CLAIM)
try:
if self.litellm_jwtauth.user_email_jwt_field is not None:
user_email = get_nested_value(
@ -460,6 +489,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:
return token.get(self.LITELLM_ORG_ID_CLAIM)
try:
if self.litellm_jwtauth.org_id_jwt_field is not None:
org_id = get_nested_value(
@ -555,53 +587,77 @@ class JWTHandler:
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=jwks_uri,
ttl=self.litellm_jwtauth.public_key_ttl,
ttl=self._get_public_key_cache_ttl(),
)
return jwks_uri
def _get_public_key_cache_ttl(self) -> float:
litellm_jwtauth = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return 600
return litellm_jwtauth.public_key_ttl
async def _get_public_key_from_jwks_url(
self, jwks_url: str, kid: Optional[str]
) -> dict:
resolved_jwks_url = await self._resolve_jwks_url(jwks_url)
cache_key = f"litellm_jwt_auth_keys_{resolved_jwks_url}"
cached_keys = await self.user_api_key_cache.async_get_cache(cache_key)
if cached_keys is None:
response = await self.http_handler.get(resolved_jwks_url)
try:
response_json = response.json()
except Exception as e:
verbose_proxy_logger.error(
f"Error parsing response: {e}. Original Response: {response.text}"
)
raise Exception(
f"Error parsing response: {e}. Check server logs for original response."
)
if "keys" in response_json:
keys: JWKKeyValue = response_json["keys"]
else:
keys = response_json
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=keys,
ttl=self._get_public_key_cache_ttl(),
)
else:
keys = cached_keys
public_key = self.parse_keys(keys=keys, kid=kid)
if public_key is not None:
return cast(dict, public_key)
raise Exception(
f"No matching public key found. keys={resolved_jwks_url}, kid={kid}"
)
async def get_public_key(self, kid: Optional[str]) -> dict:
keys_url = os.getenv("JWT_PUBLIC_KEY_URL")
if keys_url is None:
raise Exception("Missing JWT Public Key URL from environment.")
keys_url_list = [url.strip() for url in keys_url.split(",")]
keys_url_list = [url.strip() for url in keys_url.split(",") if url.strip()]
for key_url in keys_url_list:
key_url = await self._resolve_jwks_url(key_url)
cache_key = f"litellm_jwt_auth_keys_{key_url}"
cached_keys = await self.user_api_key_cache.async_get_cache(cache_key)
if cached_keys is None:
response = await self.http_handler.get(key_url)
try:
response_json = response.json()
except Exception as e:
verbose_proxy_logger.error(
f"Error parsing response: {e}. Original Response: {response.text}"
)
raise Exception(
f"Error parsing response: {e}. Check server logs for original response."
)
if "keys" in response_json:
keys: JWKKeyValue = response.json()["keys"]
else:
keys = response_json
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=keys,
ttl=self.litellm_jwtauth.public_key_ttl, # cache for 10 mins
try:
return await self._get_public_key_from_jwks_url(
jwks_url=key_url, kid=kid
)
except Exception as e:
if "No matching public key found" not in str(e):
raise
verbose_proxy_logger.debug(
"JWT Auth: No matching public key found at %s: %s", key_url, e
)
else:
keys = cached_keys
public_key = self.parse_keys(keys=keys, kid=kid)
if public_key is not None:
return cast(dict, public_key)
raise Exception(
f"No matching public key found. keys={keys_url_list}, kid={kid}"
@ -747,8 +803,157 @@ class JWTHandler:
"options": options or None,
}
def _get_configured_issuer(self, token: str) -> Optional[JWTIssuerConfig]:
litellm_jwtauth = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return None
issuer_configs = litellm_jwtauth.issuers
if not issuer_configs:
return None
claims = self.get_unverified_claims(token=token)
if claims is None:
raise Exception("Invalid JWT Submitted")
issuer = claims.get("iss")
if not isinstance(issuer, str) or not issuer:
raise Exception("JWT issuer claim is required when issuer config is set")
for issuer_config in issuer_configs:
if issuer_config.issuer == issuer:
return issuer_config
raise Exception(f"Unsupported JWT issuer: {issuer}")
def _get_jwks_url_for_issuer(self, issuer_config: JWTIssuerConfig) -> str:
if issuer_config.jwks_url:
return issuer_config.jwks_url
# _resolve_jwks_url fetches this OIDC discovery document and follows
# its jwks_uri, matching JWTIssuerConfig.jwks_url's documented fallback.
return f"{issuer_config.issuer.rstrip('/')}/.well-known/openid-configuration"
def _get_claim_value_for_issuer_mapping(
self, token: dict, claim_field: str, issuer: str
) -> Any:
sentinel = object()
claim_value = get_nested_value(
data=token,
key_path=claim_field,
default=sentinel,
)
if claim_value is sentinel or claim_value is None or claim_value == "":
raise Exception(
f"JWT issuer {issuer} missing required mapped claim: {claim_field}"
)
return claim_value
def _apply_issuer_claim_mappings(
self, token: dict, issuer_config: JWTIssuerConfig
) -> dict:
token[self.LITELLM_JWT_ISSUER_CLAIM] = issuer_config.issuer
claim_mappings = [
(issuer_config.user_id_jwt_field, self.LITELLM_USER_ID_CLAIM),
(issuer_config.user_email_jwt_field, self.LITELLM_USER_EMAIL_CLAIM),
(issuer_config.team_id_jwt_field, self.LITELLM_TEAM_ID_CLAIM),
(issuer_config.team_ids_jwt_field, self.LITELLM_TEAM_IDS_CLAIM),
(issuer_config.org_id_jwt_field, self.LITELLM_ORG_ID_CLAIM),
(issuer_config.end_user_id_jwt_field, self.LITELLM_END_USER_ID_CLAIM),
]
for source_claim, normalized_claim in claim_mappings:
if source_claim is None:
continue
token[normalized_claim] = self._get_claim_value_for_issuer_mapping(
token=token,
claim_field=source_claim,
issuer=issuer_config.issuer,
)
return token
def _get_jwk_from_public_key(self, public_key: dict) -> dict:
jwk = {}
for key in ["kty", "kid", "n", "e", "x", "y", "crv"]:
if key in public_key:
jwk[key] = public_key[key]
return jwk
def _get_decode_options(
self, audience: Optional[Union[str, List[str]]]
) -> Optional[dict]:
if audience is None:
return {"verify_aud": False}
return None
def _decode_jwt_with_public_key(
self,
token: str,
public_key: Union[dict, str],
audience: Optional[Union[str, List[str]]],
issuer: Optional[str] = None,
options: Optional[dict] = None,
) -> dict:
decode_options = (
options
if options is not None
else self._get_decode_options(audience=audience)
)
if isinstance(public_key, dict):
public_key_obj = PyJWK.from_dict(
self._get_jwk_from_public_key(public_key=public_key)
).key
return jwt.decode(
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
options=decode_options, # type: ignore[arg-type]
audience=audience,
issuer=issuer,
leeway=self.leeway,
)
cert = x509.load_pem_x509_certificate(public_key.encode(), default_backend())
key = cert.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
return jwt.decode(
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
audience=audience,
issuer=issuer,
options=decode_options, # type: ignore[arg-type]
leeway=self.leeway,
)
async def _auth_jwt_with_issuer(
self, token: str, issuer_config: JWTIssuerConfig, kid: Optional[str]
) -> dict:
public_key = await self._get_public_key_from_jwks_url(
jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config),
kid=kid,
)
try:
payload = self._decode_jwt_with_public_key(
token=token,
public_key=public_key,
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()
decode_kwargs = self._build_decode_kwargs()
header = jwt.get_unverified_header(token)
@ -756,62 +961,24 @@ class JWTHandler:
kid = header.get("kid", None)
issuer_config = self._get_configured_issuer(token=token)
if issuer_config is not None:
return await self._auth_jwt_with_issuer(
token=token,
issuer_config=issuer_config,
kid=kid,
)
public_key = await self.get_public_key(kid=kid)
if public_key is not None and isinstance(public_key, dict):
jwk = {}
if "kty" in public_key:
jwk["kty"] = public_key["kty"]
if "kid" in public_key:
jwk["kid"] = public_key["kid"]
if "n" in public_key:
jwk["n"] = public_key["n"]
if "e" in public_key:
jwk["e"] = public_key["e"]
if "x" in public_key:
jwk["x"] = public_key["x"]
if "y" in public_key:
jwk["y"] = public_key["y"]
if "crv" in public_key:
jwk["crv"] = public_key["crv"]
# parse RSA/EC/OKP keys
public_key_obj = PyJWK.from_dict(jwk).key
if public_key is not None:
try:
# decode the token using the public key
payload = jwt.decode(
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
leeway=self.leeway, # allow testing of expired tokens
**decode_kwargs,
)
return payload
except jwt.ExpiredSignatureError:
# the token is expired, do something to refresh it
raise Exception("Token Expired")
except Exception as e:
raise Exception(f"Validation fails: {str(e)}")
elif public_key is not None and isinstance(public_key, str):
try:
cert = x509.load_pem_x509_certificate(
public_key.encode(), default_backend()
)
# Extract public key
key = cert.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
# decode the token using the public key
payload = jwt.decode(
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
**decode_kwargs,
payload = self._decode_jwt_with_public_key(
token=token,
public_key=public_key,
audience=decode_kwargs["audience"],
issuer=decode_kwargs["issuer"],
options=decode_kwargs["options"],
)
return payload

View file

@ -2,6 +2,7 @@
# Unit tests for JWT-Auth
import asyncio
import base64
import os
import random
import sys
@ -21,6 +22,9 @@ from datetime import datetime, timedelta
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import jwt
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from fastapi import Request, HTTPException
from fastapi.routing import APIRoute
from fastapi.responses import Response
@ -35,7 +39,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.handle_jwt import JWTHandler, JWTAuthManager
from litellm.proxy.management_endpoints.team_endpoints import new_team
from litellm.proxy.proxy_server import chat_completion
from typing import Literal
from typing import Literal, Optional
public_key = {
"kty": "RSA",
@ -1584,3 +1588,303 @@ async def test_auth_jwt_mismatched_key_fails(monkeypatch):
with pytest.raises(Exception) as exc:
await h.auth_jwt(token)
assert "Validation fails" in str(exc.value)
def _base64url_encode_bytes(value: bytes) -> str:
return base64.urlsafe_b64encode(value).rstrip(b"=").decode()
def _base64url_encode_int(value: int) -> str:
value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big")
return _base64url_encode_bytes(value=value_bytes)
def _get_rsa_key_and_jwk(kid: str):
private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
public_numbers = private_key.public_key().public_numbers()
jwk = {
"kty": "RSA",
"n": _base64url_encode_int(value=public_numbers.n),
"e": _base64url_encode_int(value=public_numbers.e),
"kid": kid,
"alg": "RS256",
"use": "sig",
}
return private_key, jwk
def _encode_rsa_jwt(
private_key,
issuer: str,
audience: str,
kid: str,
extra_claims: Optional[dict] = None,
) -> str:
private_key_pem = private_key.private_bytes(
encoding=serialization.Encoding.PEM,
format=serialization.PrivateFormat.PKCS8,
encryption_algorithm=serialization.NoEncryption(),
)
current_time = int(time.time())
claims = {
"sub": "test-subject",
"iss": issuer,
"aud": audience,
"iat": current_time,
"exp": current_time + 300,
}
if extra_claims:
claims.update(extra_claims)
return jwt.encode(
claims,
private_key_pem,
algorithm="RS256",
headers={"kid": kid},
)
def _get_jwt_handler_with_issuer_keys(issuers: list, keys_by_url: dict) -> JWTHandler:
cache = DualCache()
for jwks_url, keys in keys_by_url.items():
cache.set_cache(
key=f"litellm_jwt_auth_keys_{jwks_url}",
value=keys,
)
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(issuers=issuers),
)
return jwt_handler
@pytest.mark.asyncio
async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims(
monkeypatch,
):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
issuer_one = "https://issuer-one.example.com"
issuer_two = "https://issuer-two.example.com"
issuer_one_jwks_url = f"{issuer_one}/keys"
issuer_two_jwks_url = f"{issuer_two}/keys"
shared_kid = "shared-kid"
_, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid)
issuer_two_private_key, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid)
jwt_handler = _get_jwt_handler_with_issuer_keys(
issuers=[
{
"issuer": issuer_one,
"jwks_url": issuer_one_jwks_url,
"audience": "audience-one",
"user_id_jwt_field": "email",
"user_email_jwt_field": "email",
},
{
"issuer": issuer_two,
"jwks_url": issuer_two_jwks_url,
"audience": "audience-two",
"user_id_jwt_field": "repository_owner",
"team_id_jwt_field": "repository",
},
],
keys_by_url={
issuer_one_jwks_url: [issuer_one_jwk],
issuer_two_jwks_url: [issuer_two_jwk],
},
)
token = _encode_rsa_jwt(
private_key=issuer_two_private_key,
issuer=issuer_two,
audience="audience-two",
kid=shared_kid,
extra_claims={
"repository_owner": "jet-ai-productivity",
"repository": "jet-ai-productivity/litellm-fork",
},
)
claims = await jwt_handler.auth_jwt(token=token)
assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two
assert jwt_handler.get_user_id(token=claims, default_value=None) == (
"jet-ai-productivity"
)
assert jwt_handler.get_team_id(token=claims, default_value=None) == (
"jet-ai-productivity/litellm-fork"
)
@pytest.mark.asyncio
async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
issuer = "https://oidc.eks.eu-west-1.amazonaws.com/id/test-cluster"
jwks_url = f"{issuer}/keys"
private_key, jwk = _get_rsa_key_and_jwk(kid="k8s-key")
jwt_handler = _get_jwt_handler_with_issuer_keys(
issuers=[
{
"issuer": issuer,
"jwks_url": jwks_url,
"audience": None,
"user_id_jwt_field": "kubernetes\\.io.namespace",
}
],
keys_by_url={jwks_url: [jwk]},
)
token = _encode_rsa_jwt(
private_key=private_key,
issuer=issuer,
audience="kubernetes.default.svc",
kid="k8s-key",
extra_claims={"kubernetes.io": {"namespace": "jet-namespace"}},
)
claims = await jwt_handler.auth_jwt(token=token)
assert jwt_handler.get_user_id(token=claims, default_value=None) == "jet-namespace"
@pytest.mark.asyncio
async def test_multi_issuer_jwt_rejects_unknown_issuer(monkeypatch):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
configured_issuer = "https://issuer.example.com"
private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key")
jwt_handler = _get_jwt_handler_with_issuer_keys(
issuers=[
{
"issuer": configured_issuer,
"jwks_url": f"{configured_issuer}/keys",
"audience": "expected-audience",
}
],
keys_by_url={f"{configured_issuer}/keys": [jwk]},
)
token = _encode_rsa_jwt(
private_key=private_key,
issuer="https://unknown-issuer.example.com",
audience="expected-audience",
kid="issuer-key",
)
with pytest.raises(Exception) as exc:
await jwt_handler.auth_jwt(token=token)
assert "Unsupported JWT issuer" in str(exc.value)
@pytest.mark.asyncio
async def test_multi_issuer_jwt_rejects_wrong_audience(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",
}
],
keys_by_url={jwks_url: [jwk]},
)
token = _encode_rsa_jwt(
private_key=private_key,
issuer=issuer,
audience="wrong-audience",
kid="issuer-key",
)
with pytest.raises(Exception) as exc:
await jwt_handler.auth_jwt(token=token)
assert "Validation fails" in str(exc.value)
@pytest.mark.asyncio
async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch):
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
issuer_one = "https://issuer-one.example.com"
issuer_two = "https://issuer-two.example.com"
issuer_one_jwks_url = f"{issuer_one}/keys"
issuer_two_jwks_url = f"{issuer_two}/keys"
shared_kid = "shared-kid"
issuer_one_private_key, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid)
_, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid)
jwt_handler = _get_jwt_handler_with_issuer_keys(
issuers=[
{
"issuer": issuer_one,
"jwks_url": issuer_one_jwks_url,
"audience": "audience-one",
},
{
"issuer": issuer_two,
"jwks_url": issuer_two_jwks_url,
"audience": "audience-two",
},
],
keys_by_url={
issuer_one_jwks_url: [issuer_one_jwk],
issuer_two_jwks_url: [issuer_two_jwk],
},
)
token = _encode_rsa_jwt(
private_key=issuer_one_private_key,
issuer=issuer_two,
audience="audience-two",
kid=shared_kid,
)
with pytest.raises(Exception) as exc:
await jwt_handler.auth_jwt(token=token)
assert "Validation fails" in str(exc.value)
@pytest.mark.asyncio
async def test_multi_issuer_jwt_missing_mapped_claim_fails_closed(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_id_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",
)
with pytest.raises(Exception) as exc:
await jwt_handler.auth_jwt(token=token)
assert "missing required mapped claim: email" in str(exc.value)