diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index b478fecaa84..a9fa10197ea 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index b7c18eb2921..957a032da5b 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -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: diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index c825bd0aab7..ca1bfdd93cb 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -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 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py index b79b38ae0ba..17fc029030b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py @@ -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(