fix(jwt,mcp): fall back to global JWKS on unknown issuer; prune fetch locks

- handle_jwt._get_configured_issuer now returns None for tokens whose 'iss'
  is not in the configured issuers list, letting auth_jwt fall through to
  the legacy JWT_PUBLIC_KEY_URL path instead of hard-raising. This keeps
  existing tokens from non-configured IdPs working when an operator adds
  the new 'issuers' list to a live deployment.

- discoverable_endpoints._prune_oauth_metadata_cache now also prunes
  entries in _OAUTH_METADATA_FETCH_LOCKS whose cache entry has been
  evicted and whose lock isn't currently held, bounding the locks dict
  to match the cache it guards.

Co-authored-by: Claude <claude@anthropic.com>
This commit is contained in:
Claude 2026-05-20 16:35:23 +00:00
parent 22e9064673
commit 488fcc25e1
No known key found for this signature in database
4 changed files with 107 additions and 14 deletions

View file

@ -53,16 +53,24 @@ def _prune_oauth_metadata_cache(now: Optional[float] = None) -> None:
for cache_key in expired_cache_keys:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
if len(_OAUTH_METADATA_CACHE) <= _OAUTH_METADATA_CACHE_MAX_SIZE:
return
if len(_OAUTH_METADATA_CACHE) > _OAUTH_METADATA_CACHE_MAX_SIZE:
overflow = len(_OAUTH_METADATA_CACHE) - _OAUTH_METADATA_CACHE_MAX_SIZE
cache_keys_by_expiry = sorted(
_OAUTH_METADATA_CACHE,
key=lambda cache_key: _OAUTH_METADATA_CACHE[cache_key][0],
)
for cache_key in cache_keys_by_expiry[:overflow]:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
overflow = len(_OAUTH_METADATA_CACHE) - _OAUTH_METADATA_CACHE_MAX_SIZE
cache_keys_by_expiry = sorted(
_OAUTH_METADATA_CACHE,
key=lambda cache_key: _OAUTH_METADATA_CACHE[cache_key][0],
)
for cache_key in cache_keys_by_expiry[:overflow]:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
# Drop locks whose cache entry has been evicted and that aren't currently
# held; held locks stay so in-flight callers continue to coalesce.
for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS):
if cache_key in _OAUTH_METADATA_CACHE:
continue
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
if lock is None or lock.locked():
continue
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
def encode_state_with_base_url(

View file

@ -883,17 +883,17 @@ class JWTHandler:
claims = self.get_unverified_claims(token=token)
if claims is None:
raise Exception("Invalid JWT Submitted")
return None
issuer = claims.get("iss")
if not isinstance(issuer, str) or not issuer:
raise Exception("JWT issuer claim is required when issuer config is set")
return None
for issuer_config in issuer_configs:
if issuer_config.issuer == issuer:
return issuer_config
raise Exception(f"Unsupported JWT issuer: {issuer}")
return None
def _get_jwks_url_for_issuer(self, issuer_config: JWTIssuerConfig) -> str:
if issuer_config.jwks_url:

View file

@ -1757,7 +1757,56 @@ async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch):
@pytest.mark.asyncio
async def test_multi_issuer_jwt_rejects_unknown_issuer(monkeypatch):
async def test_multi_issuer_jwt_falls_back_to_global_jwks_for_unknown_issuer(
monkeypatch,
):
"""Unknown ``iss`` claims fall through to the global ``JWT_PUBLIC_KEY_URL``
path so adding the new ``issuers`` config to a live deployment doesn't
break tokens minted by issuers that still rely on the legacy global JWKS.
"""
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
configured_issuer = "https://issuer.example.com"
unknown_issuer = "https://unknown-issuer.example.com"
global_jwks_url = "https://global.example.com/keys"
monkeypatch.setenv("JWT_PUBLIC_KEY_URL", global_jwks_url)
configured_private_key, configured_jwk = _get_rsa_key_and_jwk(kid="configured-key")
unknown_private_key, unknown_jwk = _get_rsa_key_and_jwk(kid="global-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": [configured_jwk],
global_jwks_url: [unknown_jwk],
},
)
token = _encode_rsa_jwt(
private_key=unknown_private_key,
issuer=unknown_issuer,
audience="expected-audience",
kid="global-key",
)
claims = await jwt_handler.auth_jwt(token=token)
assert claims["iss"] == unknown_issuer
@pytest.mark.asyncio
async def test_multi_issuer_jwt_unknown_issuer_without_global_jwks_rejected(
monkeypatch,
):
"""When there is no ``JWT_PUBLIC_KEY_URL`` to fall back to, an unknown
``iss`` claim still fails — the fallback path raises ``Missing JWT
Public Key URL`` rather than the legacy ``Unsupported JWT issuer``.
"""
monkeypatch.delenv("JWT_AUDIENCE", raising=False)
monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False)
@ -1783,7 +1832,7 @@ async def test_multi_issuer_jwt_rejects_unknown_issuer(monkeypatch):
with pytest.raises(Exception) as exc:
await jwt_handler.auth_jwt(token=token)
assert "Unsupported JWT issuer" in str(exc.value)
assert "Missing JWT Public Key URL" in str(exc.value)
@pytest.mark.asyncio

View file

@ -7,7 +7,9 @@ Covers:
network errors as HTTP 502).
"""
import asyncio
import sys
import time
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -20,6 +22,7 @@ sys.path.insert(0, "../../../../../")
from litellm.proxy._experimental.mcp_server import discoverable_endpoints
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
_OAUTH_METADATA_CACHE,
_OAUTH_METADATA_FETCH_LOCKS,
_build_oauth_protected_resource_response,
)
from litellm.proxy._types import MCPTransport
@ -42,8 +45,10 @@ def _mock_mcp_client_ip():
def _clear_metadata_cache():
"""Prevent cross-test cache bleed for the oauth-protected-resource TTL cache."""
_OAUTH_METADATA_CACHE.clear()
_OAUTH_METADATA_FETCH_LOCKS.clear()
yield
_OAUTH_METADATA_CACHE.clear()
_OAUTH_METADATA_FETCH_LOCKS.clear()
def _make_request(base_url: str = "https://gateway.example.com/") -> Request:
@ -233,6 +238,37 @@ def test_oauth_metadata_cache_prunes_to_max_size():
) in _OAUTH_METADATA_CACHE
def test_oauth_metadata_fetch_locks_pruned_alongside_cache():
now = 1_000_000.0
cached_key = ("server-active", "https://upstream/active")
expired_key = ("server-expired", "https://upstream/expired")
orphan_key = ("server-orphan", "https://upstream/orphan")
_OAUTH_METADATA_CACHE[cached_key] = (now + 100, {"index": 0})
_OAUTH_METADATA_CACHE[expired_key] = (now - 1, {"index": 1})
_OAUTH_METADATA_FETCH_LOCKS[cached_key] = asyncio.Lock()
_OAUTH_METADATA_FETCH_LOCKS[expired_key] = asyncio.Lock()
_OAUTH_METADATA_FETCH_LOCKS[orphan_key] = asyncio.Lock()
discoverable_endpoints._prune_oauth_metadata_cache(now)
assert cached_key in _OAUTH_METADATA_FETCH_LOCKS
assert expired_key not in _OAUTH_METADATA_FETCH_LOCKS
assert orphan_key not in _OAUTH_METADATA_FETCH_LOCKS
@pytest.mark.asyncio
async def test_oauth_metadata_fetch_locks_held_lock_preserved_during_prune():
held_key = ("server-busy", "https://upstream/busy")
held_lock = asyncio.Lock()
_OAUTH_METADATA_FETCH_LOCKS[held_key] = held_lock
async with held_lock:
discoverable_endpoints._prune_oauth_metadata_cache(time.time())
assert held_key in _OAUTH_METADATA_FETCH_LOCKS
@pytest.mark.asyncio
async def test_oauth_metadata_cache_expired_entry_is_refetched():
passthrough_server = MCPServer(