mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
22e9064673
commit
488fcc25e1
4 changed files with 107 additions and 14 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue