test(proxy): cover issuer-scoped JWT auth

This commit is contained in:
gym-cmd 2026-05-20 23:52:51 +01:00
parent 96e4de3c48
commit 2636bbcdc7

View file

@ -2754,3 +2754,564 @@ def test_build_decode_kwargs_no_warning_when_scoped(
if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage()
]
assert matching == []
def _base64url_encode_int(value: int) -> str:
import base64
value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big")
return base64.urlsafe_b64encode(value_bytes).decode("utf-8").rstrip("=")
def _get_rsa_key_and_jwk(kid: str):
from cryptography.hazmat.primitives.asymmetric import rsa
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:
import time
import jwt
from cryptography.hazmat.primitives import serialization
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:
from litellm.caching.dual_cache import DualCache
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_get_public_key_fetches_and_caches_jwks_response():
from unittest.mock import AsyncMock, MagicMock
from litellm.caching.dual_cache import DualCache
jwt_handler = JWTHandler()
cache = DualCache()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(public_key_ttl=123),
)
expected_key_id = "cached-key"
_, jwk = _get_rsa_key_and_jwk(kid=expected_key_id)
mock_response = MagicMock()
mock_response.json.return_value = {"keys": [jwk]}
jwt_handler.http_handler.get = AsyncMock(return_value=mock_response)
public_key = await jwt_handler._get_public_key_from_jwks_url(
jwks_url="https://issuer.example.com/keys",
kid=expected_key_id,
)
assert public_key == jwk
cached_keys = await cache.async_get_cache(
key="litellm_jwt_auth_keys_https://issuer.example.com/keys"
)
assert cached_keys == [jwk]
@pytest.mark.asyncio
async def test_get_public_key_tries_next_jwks_url_when_kid_missing(monkeypatch):
from litellm.caching.dual_cache import DualCache
first_jwks_url = "https://first.example.com/keys"
second_jwks_url = "https://second.example.com/keys"
monkeypatch.setenv(
"JWT_PUBLIC_KEY_URL", f"{first_jwks_url}, {second_jwks_url},,"
)
_, first_jwk = _get_rsa_key_and_jwk(kid="first-key")
_, second_jwk = _get_rsa_key_and_jwk(kid="second-key")
cache = DualCache()
cache.set_cache(key=f"litellm_jwt_auth_keys_{first_jwks_url}", value=[first_jwk])
cache.set_cache(
key=f"litellm_jwt_auth_keys_{second_jwks_url}", value=[second_jwk]
)
jwt_handler = JWTHandler()
jwt_handler.update_environment(
prisma_client=None,
user_api_key_cache=cache,
litellm_jwtauth=LiteLLM_JWTAuth(),
)
public_key = await jwt_handler.get_public_key(kid="second-key")
assert public_key == second_jwk
def test_get_jwks_url_for_issuer_falls_back_to_discovery_document():
jwt_handler = JWTHandler()
issuer_config = LiteLLM_JWTAuth(
issuers=[{"issuer": "https://issuer.example.com/tenant/"}]
).issuers[0]
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"
)
@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": "example-org",
"repository": "example-org/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) == "example-org"
assert jwt_handler.get_team_id(token=claims, default_value=None) == (
"example-org/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,
"disable_audience_validation": True,
"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": "example-namespace"}},
)
claims = await jwt_handler.auth_jwt(token=token)
assert (
jwt_handler.get_user_id(token=claims, default_value=None) == "example-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)
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):
from litellm.caching.dual_cache import DualCache
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
):
import logging
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