From 444894be2ad0493600de50e673593ba8bf04597c Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 20:01:28 +0000 Subject: [PATCH 1/7] fix(mcp): key discovery caches per caller correctly and drop stale caches on server updates Discovery-list cache identity now uses the hashed token instead of the raw api_key and treats MCPJWTSigner-signed servers as per caller. Server definition changes also drop the cached upstream OAuth metadata. OpenAPI listings look tools up under the normalized registry prefix with the separator, so an overlapping sibling prefix no longer leaks into the list. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 11 +++ .../mcp_server/mcp_server_manager.py | 70 ++++++++------ .../mcp_server/test_discoverable_endpoints.py | 35 +++++++ .../mcp_server/test_mcp_server_manager.py | 92 +++++++++++++++++++ 4 files changed, 180 insertions(+), 28 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d42c1c6b879..91da7204435 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -141,6 +141,17 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) +def invalidate_oauth_metadata_cache(server_id: str) -> None: + """Drop cached upstream IdP metadata for a server whose definition changed.""" + for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: + del _OAUTH_METADATA_CACHE[cache_key] + for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: + 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( base_url: str, original_state: str, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 31896d9ddc5..d2a4b6ebad9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2664,7 +2664,7 @@ class MCPServerManager: self._assign_unique_short_prefix(new_server) _warn_legacy_delegate_auth_if_applicable(new_server, source="config") _warn_config_id_jag_server_outruns_sso(new_server) - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.config_mcp_servers[server_id] = new_server self._set_oauth_discovery_deferred( server_id, @@ -2868,7 +2868,7 @@ class MCPServerManager: global_mcp_tool_registry, ) - self._invalidate_discovery_lists(server.server_id) + self._invalidate_server_definition_caches(server.server_id) prefix_root: Final = normalize_server_name(get_server_prefix(server)) if server.spec_path and prefix_root: openapi_key_prefix: Final = prefix_root + MCP_TOOL_PREFIX_SEPARATOR @@ -3275,7 +3275,7 @@ class MCPServerManager: # env_vars_are_encrypted=False. new_server: Final = await self.build_mcp_server_from_table(mcp_server, env_vars_are_encrypted=False) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -3312,7 +3312,7 @@ class MCPServerManager: previous_server=self.registry[mcp_server.server_id], ) self._assign_unique_short_prefix(new_server) - self._invalidate_discovery_lists(mcp_server.server_id) + self._invalidate_server_definition_caches(mcp_server.server_id) self.registry[mcp_server.server_id] = new_server await self._maybe_register_openapi_tools(new_server) self.prime_oauth_metadata_discovery(new_server) @@ -4492,24 +4492,16 @@ class MCPServerManager: if server.spec_path: # OpenAPI tools were stored in the registry under the prefix # active at registration time — fetch by that same prefix. - _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=get_server_prefix(server)) + registry_prefix: Final = normalize_server_name(get_server_prefix(server)) + MCP_TOOL_PREFIX_SEPARATOR + _tools: Final = global_mcp_tool_registry.list_tools(tool_prefix=registry_prefix) tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(_tools) # OpenAPI tools are stored in the registry with their prefix already # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". - if not add_prefix: - prefix: Final = get_server_prefix(server) - sep: Final = MCP_TOOL_PREFIX_SEPARATOR - tools = [ - ( - t.model_copy(update={"name": t.name[len(prefix) + len(sep) :]}) - if t.name.startswith(f"{prefix}{sep}") - else t - ) - for t in tools - ] - return tools + if add_prefix: + return tools + return [t.model_copy(update=MappingProxyType({"name": t.name[len(registry_prefix) :]})) for t in tools] else: tools = await self._fetch_tools_with_timeout(client, server.name) self._remember_upstream_initialize_instructions(server, client) @@ -4558,6 +4550,35 @@ class MCPServerManager: self._resource_discovery_cache.invalidate(server_id) self._template_discovery_cache.invalidate(server_id) + def _invalidate_server_definition_caches(self, server_id: str) -> None: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # lazy: discoverable_endpoints lazily imports this module's manager singleton + invalidate_oauth_metadata_cache, + ) + + self._invalidate_discovery_lists(server_id) + invalidate_oauth_metadata_cache(server_id) + + def _discovers_per_caller(self, server: MCPServer) -> bool: + return ( + server.requires_per_user_auth + or self._references_per_user_env_var(server) + or server.delegate_auth_to_upstream + or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + or self._signs_caller_identity_upstream(server) + ) + + @staticmethod + def _signs_caller_identity_upstream(server: MCPServer) -> bool: + """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may + tailor its catalog to the caller even though the server itself is configured as shared.""" + from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server + get_mcp_jwt_signer, + ) + + if get_mcp_jwt_signer() is None: + return False + return not any(k.lower() == "authorization" for k in (server.static_headers or {})) + def _discovery_key( self, server: MCPServer, @@ -4568,25 +4589,18 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = ( - server.requires_per_user_auth - or self._references_per_user_env_var(server) - or server.delegate_auth_to_upstream - or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) - ) + per_user: Final = self._discovers_per_caller(server) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.api_key) - if per_user and user_api_key_auth is not None - else None + (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), sort_keys=True, separators=(",", ":"), ) - return server.server_id, hashlib.sha256(material.encode()).hexdigest() + return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() async def get_prompts_from_server( self, @@ -6704,7 +6718,7 @@ class MCPServerManager: for server_id in previous_registry.keys() | registered_registry.keys(): if previous_registry.get(server_id) != registered_registry.get(server_id): - self._invalidate_discovery_lists(server_id) + self._invalidate_server_definition_caches(server_id) self.registry = registered_registry _warn_on_shared_identifier_prefixes(registered_registry.values()) # A discovery task may have published into ``previous_registry`` while diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index f9a0075e530..0a92b548d62 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12611,3 +12611,38 @@ async def test_identity_bound_authorize_unrelated_bearer_uses_browser_session( proxy_server.prisma_client.db.litellm_mcpusercredentials.upsert.assert_not_called() proxy_server.prisma_client.db.litellm_usertable.create.assert_not_called() proxy_server.prisma_client.db.litellm_teamtable.create.assert_not_called() + + +@pytest.mark.asyncio +async def test_update_server_drops_cached_upstream_oauth_metadata(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._types import LiteLLM_MCPServerTable, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + server = MCPServer( + server_id="oauth-cache-server", + name="oauth_cache_server", + url="http://old-upstream/mcp", + transport=MCPTransport.http, + ) + manager.registry[server.server_id] = server + stale_key: Final = (server.server_id, server.url) + other_key: Final = ("other-server", "http://other/mcp") + discoverable_endpoints._OAUTH_METADATA_CACHE[stale_key] = (time.time() + 300, {"iss": "old-idp"}) + discoverable_endpoints._OAUTH_METADATA_CACHE[other_key] = (time.time() + 300, {"iss": "other"}) + try: + await manager.update_server( + LiteLLM_MCPServerTable( + server_id=server.server_id, + server_name=server.name, + url="http://new-upstream/mcp", + transport=MCPTransport.http, + ) + ) + assert stale_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert other_key in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 70ef4312f4c..85b4aaa11eb 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14058,6 +14058,98 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) + from litellm.proxy._types import hash_token + + same_user_other_token: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None + ) + same_token_no_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None + ) + with_key: Final = manager._discovery_key( + server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None + ) + assert same_user_other_token != with_key + assert same_token_no_key == with_key + + +@pytest.mark.parametrize( + ("signer", "static_headers", "shared"), + [ + pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), + pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), + pytest.param(None, None, True, id="no-signer-stays-shared"), + ], +) +def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None: + manager: Final = MCPServerManager() + server: Final = _discovery_server().model_copy(update={"static_headers": static_headers}) + alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice") + bob: Final = UserAPIKeyAuth(user_id="bob", token="hashed-bob") + + with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam + "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", + return_value=signer, + ): + for_alice: Final = manager._discovery_key(server, alice, None, None, None, None) + for_bob: Final = manager._discovery_key(server, bob, None, None, None, None) + + assert (for_alice == for_bob) is shared + + +def _register_local_tool(name: str, description: str) -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + async def _handler(**kwargs): + return None + + global_mcp_tool_registry.register_tool( + name=name, description=description, input_schema={"type": "object"}, handler=_handler + ) + + +def _openapi_server(name: str) -> MCPServer: + return MCPServer( + server_id=f"{name}-id", name=name, alias=name, transport=MCPTransport.http, url=None, spec_path="/spec.yaml" + ) + + +@pytest.mark.asyncio +async def test_openapi_listing_ignores_overlapping_server_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + _register_local_tool("pet-list", "Local pet tool") + _register_local_tool("petstore-list", "Foreign petstore tool") + try: + prefixed: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=True) + bare: Final = await manager._get_tools_from_server(server=_openapi_server("pet"), add_prefix=False) + finally: + for prefix in ("pet-", "petstore-"): + global_mcp_tool_registry.unregister_tools_with_prefix(prefix) + + assert [t.name for t in prefixed] == ["pet-list"] + assert [t.name for t in bare] == ["list"] + + +@pytest.mark.asyncio +async def test_openapi_listing_finds_tools_registered_under_the_normalized_prefix() -> None: + from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry + + manager: Final = MCPServerManager() + manager._create_mcp_client = AsyncMock(return_value=AsyncMock()) + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + _register_local_tool("pet_store-list", "Pet store tool") + try: + listed: Final = await manager._get_tools_from_server(server=_openapi_server("pet store"), add_prefix=False) + finally: + global_mcp_tool_registry.unregister_tools_with_prefix("pet_store-") + + assert [t.name for t in listed] == ["list"] + @pytest.mark.asyncio async def test_discovery_cache_retries_cancelled_fetches() -> None: From cee44351e306c75be29b739544e71d3a5209f9b9 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 20:17:43 +0000 Subject: [PATCH 2/7] fix(mcp): keep the discovery cache digest call unchanged so CodeQL matches the existing alert Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/_experimental/mcp_server/mcp_server_manager.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d2a4b6ebad9..4504f4a85d2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4600,7 +4600,7 @@ class MCPServerManager: sort_keys=True, separators=(",", ":"), ) - return server.server_id, hashlib.sha256(material.encode(), usedforsecurity=False).hexdigest() + return server.server_id, hashlib.sha256(material.encode()).hexdigest() async def get_prompts_from_server( self, From 9ce2803f429d2370fd378c5a8d4cbd30f4ad7791 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:03:22 +0000 Subject: [PATCH 3/7] fix(mcp): guard OAuth metadata cache writes with a per-server generation and drop unproven per-caller discovery keys An upstream metadata fetch that started before a server edit could store its stale reply after invalidate_oauth_metadata_cache ran. Invalidation now bumps a per-server generation and the fetch only stores when the generation it captured before I/O is unchanged. The MCPJWTSigner-based per-caller discovery classification and the api_key to token key change had no reproduction (the signer only injects on tools/list, and UserAPIKeyAuth hashes api_key in place), so both go back to the merge-base behavior. Integration coverage under tests/integration/mcp: overlapping OpenAPI aliases, a config-declared server name with a space, OAuth metadata refetch after a save, and the in-flight stale-write race Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 26 ++-- .../mcp_server/mcp_server_manager.py | 32 ++--- tests/integration/mcp/test_mcp_management.py | 52 ++++++++ .../mcp/test_oauth_configuration.py | 115 +++++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 43 +++++++ .../mcp_server/test_mcp_server_manager.py | 38 ------ 6 files changed, 228 insertions(+), 78 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 91da7204435..d2851cf4fbf 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -107,6 +107,9 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Per-server_id generation, bumped on invalidation so a fetch that started before the server +# definition changed cannot repopulate the cache with the stale reply. +_OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( tags=["mcp"], @@ -143,6 +146,7 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: @@ -2377,6 +2381,14 @@ async def fetch_upstream_oauth_protected_resource( cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: return cached[1] + generation: Final = _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) + + def store(payload: dict | None, ttl_seconds: int) -> None: + if _OAUTH_METADATA_GENERATIONS.get(mcp_server.server_id, 0) != generation: + return + stored_at: Final = time.time() + _OAUTH_METADATA_CACHE[cache_key] = (stored_at + ttl_seconds, payload) + _prune_oauth_metadata_cache(stored_at) host_base: Final = f"{upstream.scheme}://{upstream.netloc}" candidates: Final = [f"{host_base}/.well-known/oauth-protected-resource"] @@ -2418,12 +2430,7 @@ async def fetch_upstream_oauth_protected_resource( ) continue if isinstance(payload, dict): - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_CACHE_TTL_SECONDS, - payload, - ) - _prune_oauth_metadata_cache(now) + store(payload, _OAUTH_METADATA_CACHE_TTL_SECONDS) return payload if len(network_errors) == len(candidates): @@ -2432,12 +2439,7 @@ async def fetch_upstream_oauth_protected_resource( # Negative-result caching: when no candidate yielded a usable payload, # remember that for a shorter TTL so we don't re-fetch on every # subsequent discovery request (and so the per-key lock can be pruned). - now = time.time() - _OAUTH_METADATA_CACHE[cache_key] = ( - now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS, - None, - ) - _prune_oauth_metadata_cache(now) + store(None, _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS) return None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 4504f4a85d2..ca05266c827 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -4558,27 +4558,6 @@ class MCPServerManager: self._invalidate_discovery_lists(server_id) invalidate_oauth_metadata_cache(server_id) - def _discovers_per_caller(self, server: MCPServer) -> bool: - return ( - server.requires_per_user_auth - or self._references_per_user_env_var(server) - or server.delegate_auth_to_upstream - or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) - or self._signs_caller_identity_upstream(server) - ) - - @staticmethod - def _signs_caller_identity_upstream(server: MCPServer) -> bool: - """Whether MCPJWTSigner mints a per-caller ``Authorization`` for ``server``, so the upstream may - tailor its catalog to the caller even though the server itself is configured as shared.""" - from litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer import ( # noqa: PLC0415 # lazy: guardrail package imports the proxy server - get_mcp_jwt_signer, - ) - - if get_mcp_jwt_signer() is None: - return False - return not any(k.lower() == "authorization" for k in (server.static_headers or {})) - def _discovery_key( self, server: MCPServer, @@ -4589,11 +4568,18 @@ class MCPServerManager: subject_token: str | None, credential_fingerprint: str | None = None, ) -> _DiscoveryKey: - per_user: Final = self._discovers_per_caller(server) + per_user: Final = ( + server.requires_per_user_auth + or self._references_per_user_env_var(server) + or server.delegate_auth_to_upstream + or server.auth_type in (MCPAuth.oauth2_token_exchange, MCPAuth.oauth2_id_jag) + ) if not (per_user or mcp_auth_header or extra_headers or stdio_env or subject_token): return server.server_id, None identity: Final = ( - (user_api_key_auth.user_id, user_api_key_auth.token) if per_user and user_api_key_auth is not None else None + (user_api_key_auth.user_id, user_api_key_auth.api_key) + if per_user and user_api_key_auth is not None + else None ) material: Final = json.dumps( (identity, mcp_auth_header, extra_headers, stdio_env, subject_token, credential_fingerprint), diff --git a/tests/integration/mcp/test_mcp_management.py b/tests/integration/mcp/test_mcp_management.py index 66c30a62bde..67cdbbff5a4 100644 --- a/tests/integration/mcp/test_mcp_management.py +++ b/tests/integration/mcp/test_mcp_management.py @@ -1,3 +1,4 @@ +import itertools import uuid from pathlib import Path from typing import Final @@ -7,10 +8,13 @@ import yaml from integration._support.client import Gateway, eventually from integration._support.mcp import ( McpCaller, + McpPeer, call_tool, delete_mcp, forget_mcp, + listed_tools, mcp_peer, + openapi_peer, register_mcp, tool_calls, tool_names, @@ -189,6 +193,54 @@ def test_duplicate_alias_is_rejected_so_tool_prefixes_cannot_collide(gateway: Ga scenario.cleanups.callback(forget_mcp, gateway, winner) +def _openapi_server_lists_and_calls_only_its_own_tools( + gateway: Gateway, key: str, peer: McpPeer, identity: str +) -> None: + listed: Final = set(listed_tools(gateway, key, identity)) + assert listed == {"getpet", "createpet"}, (identity, listed) + peer.drain() + called: Final = call_tool(gateway, key, identity, "getpet", {"petId": "7"}) + assert called.status_code == 200, called.text + assert [(item["method"], item["path"]) for item in peer.drain()] == [("GET", "/pets/7")], identity + + +def test_openapi_listing_is_scoped_to_the_exact_alias_when_aliases_overlap(gateway: Gateway) -> None: + with openapi_peer() as short, openapi_peer() as long, gateway.scenario() as scenario: + stem: Final = "pet" + uuid.uuid4().hex[:8] + servers: Final = tuple( + (peer, alias, register_mcp(scenario, peer, alias)) + for peer, alias in ((short, stem), (long, stem + "store")) + ) + key: Final = scenario.key(object_permission={"mcp_servers": [identity for _, _, identity in servers]}) + for peer, _, identity in servers: + _openapi_server_lists_and_calls_only_its_own_tools(gateway, key, peer, identity) + aggregate: Final = McpCaller(gateway, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + assert sorted(aggregate.tools) == sorted( + f"{prefix}-{tool}" for prefix, tool in itertools.product((stem, stem + "store"), ("getpet", "createpet")) + ), aggregate.tools + assert all(peer.drain() == () for peer, _, _ in servers), "listing must not reach any OpenAPI upstream" + + +def test_config_declared_openapi_server_with_a_space_in_its_name_lists_its_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with openapi_peer() as peer: + config: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + name: Final = "pet store " + uuid.uuid4().hex[:8] + config["mcp_servers"] = {name: peer.registration()} + path: Final = tmp_path / "openapi-space.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + identity: Final = next(i for i, s in _servers(candidate).items() if s["server_name"] == name) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + _openapi_server_lists_and_calls_only_its_own_tools(candidate, key, peer, identity) + aggregate: Final = McpCaller(candidate, key, "mcp").list_tools() + assert aggregate.ok, aggregate.raw + prefix: Final = name.replace(" ", "_") + assert sorted(aggregate.tools) == [f"{prefix}-createpet", f"{prefix}-getpet"], aggregate.tools + + def test_invalid_registrations_are_rejected(gateway: Gateway) -> None: with mcp_peer() as peer, gateway.scenario() as scenario: alias: Final = "mgmt" + uuid.uuid4().hex[:8] diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 4c46c706054..3a023f6d38e 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -1,17 +1,23 @@ import json import queue +import threading import uuid -from urllib.parse import parse_qs, urlsplit -from typing import Final, Literal +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field from pathlib import Path +from typing import Final, Literal +from urllib.parse import parse_qs, urlsplit import pytest - -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, Scenario, eventually from integration._support.database import read_rows from integration._support.mcp import McpPeer, call_tool, mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy -from integration._support.wire import Reply, Request, wire_server +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import TypeAdapter + +_Upstream = Callable[[Request], Reply] @pytest.mark.covers("other.mcp.oauth.discovery_cannot_erase_configured_authorization_endpoint") @@ -104,6 +110,105 @@ def test_partial_discovery_and_unrelated_edit_keep_actual_authorization_destinat assert updated.status_code == 202, updated.text +@dataclass(frozen=True, slots=True) +class _Hold: + armed: threading.Event = field(default_factory=threading.Event) + released: threading.Event = field(default_factory=threading.Event) + + +def _idp_upstream(origin: Callable[[], str], moved: threading.Event, hold: _Hold | None = None) -> _Upstream: + def issuer() -> str: + return origin() + ("/idp-after" if moved.is_set() else "/idp-before") + + def respond(request: Request) -> Reply: + if "oauth-authorization-server" in request.target or "openid-configuration" in request.target: + current: Final = issuer() + return Reply( + body=json.dumps( + { + "issuer": current, + "authorization_endpoint": current + "/authorize", + "token_endpoint": current + "/token", + } + ).encode() + ) + if request.target.startswith("/.well-known/oauth-protected-resource"): + body: Final = json.dumps({"resource": origin() + "/mcp", "authorization_servers": [issuer()]}).encode() + if hold is not None and hold.armed.is_set(): + assert hold.released.wait(timeout=15), "the held upstream metadata reply was never released" + return Reply(body=body) + return Reply(status=404, body=b'{"error":"unexpected"}') + + return respond + + +def _register_pass_through(scenario: Scenario, wire: Wire, alias: str) -> str: + return register_mcp(scenario, McpPeer(wire.url + "/mcp", queue.Queue()), alias, auth_type="true_passthrough") + + +def _wire_requests(wire: Wire, seen: list[Request]) -> Callable[[], tuple[Request, ...]]: + def observed() -> tuple[Request, ...]: + seen.extend(wire.drain()) + return tuple(seen) + + return observed + + +def _registration_discovery_settled(requests: tuple[Request, ...]) -> bool: + return any( + "oauth-authorization-server" in item.target or "openid-configuration" in item.target for item in requests + ) + + +def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str, ...]: + response: Final = gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp") + assert response.status_code == 200, response.text + return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) + + +def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + moved.set() + wire.drain() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + assert any(request.target.startswith("/.well-known/oauth-protected-resource") for request in wire.drain()), ( + "the save must send protected-resource discovery back to the upstream" + ) + + +def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: + moved: Final = threading.Event() + hold: Final = _Hold() + with ( + wire_server(_idp_upstream(lambda: wire.url, moved, hold)) as wire, + gateway.scenario() as scenario, + ThreadPoolExecutor(max_workers=1) as pool, + ): + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + seen: Final[list[Request]] = [] + observed: Final = _wire_requests(wire, seen) + eventually(observed, _registration_discovery_settled, seconds=10) + settled: Final = len(seen) + hold.armed.set() + stale: Final = pool.submit(_advertised_authorization_servers, gateway, alias) + eventually(observed, lambda requests: len(requests) > settled, seconds=10) + assert seen[settled].target.startswith("/.well-known/oauth-protected-resource"), seen[settled:] + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + hold.released.set() + assert stale.result(timeout=30) == (wire.url + "/idp-before",) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + + @pytest.mark.covers("other.mcp.oauth.same_url_credentials_are_isolated_by_user_and_server") @pytest.mark.parametrize("transition", ("revoke", "expire")) def test_same_url_oauth_credentials_and_revocation_are_isolated_by_user_and_server( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0a92b548d62..7a2af8dfcca 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12646,3 +12646,46 @@ async def test_update_server_drops_cached_upstream_oauth_metadata(): finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(stale_key, None) discoverable_endpoints._OAUTH_METADATA_CACHE.pop(other_key, None) + + +@pytest.mark.asyncio +async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cache(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="stale-write-server", name="stale_write", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["old-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + in_flight: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await in_flight == {"authorization_servers": ["old-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 85b4aaa11eb..130d47bbafa 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -14058,44 +14058,6 @@ def test_discovery_cache_keys_isolate_user_dependent_auth(auth_type: MCPAuth) -> assert "first" not in str(first) assert "second" not in str(second) - from litellm.proxy._types import hash_token - - same_user_other_token: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-second")), None, None, None, None - ) - same_token_no_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", token=hash_token("sk-first")), None, None, None, None - ) - with_key: Final = manager._discovery_key( - server, UserAPIKeyAuth(user_id="first", api_key="sk-first"), None, None, None, None - ) - assert same_user_other_token != with_key - assert same_token_no_key == with_key - - -@pytest.mark.parametrize( - ("signer", "static_headers", "shared"), - [ - pytest.param(MagicMock(), None, False, id="signer-mints-per-caller-authorization"), - pytest.param(MagicMock(), {"Authorization": "Bearer admin-token"}, True, id="static-authorization-wins"), - pytest.param(None, None, True, id="no-signer-stays-shared"), - ], -) -def test_jwt_signer_makes_a_shared_server_discover_per_caller(signer, static_headers, shared) -> None: - manager: Final = MCPServerManager() - server: Final = _discovery_server().model_copy(update={"static_headers": static_headers}) - alice: Final = UserAPIKeyAuth(user_id="alice", token="hashed-alice") - bob: Final = UserAPIKeyAuth(user_id="bob", token="hashed-bob") - - with patch( # test-quality-ok: the signer is a process-wide singleton the manager reads, no injection seam - "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", - return_value=signer, - ): - for_alice: Final = manager._discovery_key(server, alice, None, None, None, None) - for_bob: Final = manager._discovery_key(server, bob, None, None, None, None) - - assert (for_alice == for_bob) is shared - def _register_local_tool(name: str, description: str) -> None: from litellm.proxy._experimental.mcp_server.tool_registry import global_mcp_tool_registry From c0099a45deb9b4c122c95a6b768e6c3583964e9a Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:37:53 +0000 Subject: [PATCH 4/7] fix(mcp): keep OAuth metadata generations only while a fetch is in flight Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 15 +++++++++++++-- .../mcp_server/test_discoverable_endpoints.py | 16 ++++++++++++++++ 2 files changed, 29 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d2851cf4fbf..d4e95e7ad70 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -108,7 +108,8 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} # Per-server_id generation, bumped on invalidation so a fetch that started before the server -# definition changed cannot repopulate the cache with the stale reply. +# definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch +# in flight carry an entry; the rest are pruned with the cache. _OAUTH_METADATA_GENERATIONS: Final[dict[str, int]] = {} router: Final = APIRouter( @@ -143,10 +144,20 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + for server_id in [sid for sid in _OAUTH_METADATA_GENERATIONS if not _oauth_metadata_fetch_in_flight(sid)]: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) + + +def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: + return any(lock.locked() for cache_key, lock in _OAUTH_METADATA_FETCH_LOCKS.items() if cache_key[0] == server_id) + def invalidate_oauth_metadata_cache(server_id: str) -> None: """Drop cached upstream IdP metadata for a server whose definition changed.""" - _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + if _oauth_metadata_fetch_in_flight(server_id): + _OAUTH_METADATA_GENERATIONS[server_id] = _OAUTH_METADATA_GENERATIONS.get(server_id, 0) + 1 + else: + _OAUTH_METADATA_GENERATIONS.pop(server_id, None) for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 7a2af8dfcca..d9135a999b6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12686,6 +12686,22 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach release.set() assert await in_flight == {"authorization_servers": ["old-idp"]} assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + discoverable_endpoints._prune_oauth_metadata_cache() + assert server.server_id not in discoverable_endpoints._OAUTH_METADATA_GENERATIONS finally: discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + +def test_invalidating_an_idle_server_leaves_no_generation_behind(): + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache + + server_ids: Final = tuple(f"churned-server-{i}" for i in range(50)) + try: + for server_id in server_ids: + invalidate_oauth_metadata_cache(server_id) + assert not set(server_ids) & set(discoverable_endpoints._OAUTH_METADATA_GENERATIONS) + finally: + for server_id in server_ids: + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server_id, None) From 9786e3509aee0ab937c85946c91e084f3513918b Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 22:58:32 +0000 Subject: [PATCH 5/7] fix(mcp): count queued OAuth metadata fetchers so invalidation survives lock handoff Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 38 ++++++++----- .../mcp_server/test_discoverable_endpoints.py | 54 +++++++++++++++++++ 2 files changed, 80 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index d4e95e7ad70..dd72328732b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -3,7 +3,8 @@ import html as _html import json import secrets import time -from collections.abc import Callable, Mapping +from collections.abc import AsyncIterator, Callable, Mapping +from contextlib import asynccontextmanager from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Final, Literal, Optional from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse @@ -107,6 +108,10 @@ _OAUTH_METADATA_CACHE_MAX_SIZE: Final = 128 # Per-(server_id, resource_url) async locks so concurrent discovery requests # coalesce onto a single upstream fetch instead of issuing N parallel calls. _OAUTH_METADATA_FETCH_LOCKS: Final[dict[tuple[str, str], asyncio.Lock]] = {} +# Callers inside ``_oauth_metadata_fetch_slot`` per cache key, lock waiters included. ``Lock.locked()`` +# reads False between one holder's release and the next waiter's wake-up, so it cannot tell an +# idle lock from one being handed off. +_OAUTH_METADATA_FETCHERS: Final[dict[tuple[str, str], int]] = {} # Per-server_id generation, bumped on invalidation so a fetch that started before the server # definition changed cannot repopulate the cache with the stale reply. Only servers with a fetch # in flight carry an entry; the rest are pruned with the cache. @@ -134,13 +139,10 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: 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. + # Drop locks whose cache entry has been evicted and that nobody holds or + # waits on; the rest 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(): + if cache_key in _OAUTH_METADATA_CACHE or cache_key in _OAUTH_METADATA_FETCHERS: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -149,7 +151,21 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: - return any(lock.locked() for cache_key, lock in _OAUTH_METADATA_FETCH_LOCKS.items() if cache_key[0] == server_id) + return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) + + +@asynccontextmanager +async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]: + _OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1 + try: + async with _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()): + yield + finally: + remaining: Final = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) - 1 + if remaining > 0: + _OAUTH_METADATA_FETCHERS[cache_key] = remaining + else: + _OAUTH_METADATA_FETCHERS.pop(cache_key, None) def invalidate_oauth_metadata_cache(server_id: str) -> None: @@ -161,8 +177,7 @@ def invalidate_oauth_metadata_cache(server_id: str) -> None: for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: - lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) - if lock is None or lock.locked(): + if cache_key in _OAUTH_METADATA_FETCHERS: continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -2386,8 +2401,7 @@ async def fetch_upstream_oauth_protected_resource( if cached is not None and cached[0] > now: return cached[1] - lock: Final = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock()) - async with lock: + async with _oauth_metadata_fetch_slot(cache_key): now = time.time() cached = _OAUTH_METADATA_CACHE.get(cache_key) if cached is not None and cached[0] > now: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index d9135a999b6..b1f0b3fa67e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -12693,6 +12693,60 @@ async def test_metadata_fetched_before_invalidation_does_not_repopulate_the_cach discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) +@pytest.mark.asyncio +async def test_fetch_waiting_on_a_lock_handoff_stays_tracked_through_invalidation(): + import asyncio + + from litellm.proxy._experimental.mcp_server import discoverable_endpoints + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + fetch_upstream_oauth_protected_resource, + invalidate_oauth_metadata_cache, + ) + from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="handoff-server", name="handoff", url="http://upstream/mcp", transport=MCPTransport.http + ) + cache_key: Final = (server.server_id, server.url) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def slow_get(url: str, headers: dict[str, str]) -> MagicMock: + started.set() + await release.wait() + return MagicMock(status_code=200, json=MagicMock(return_value={"authorization_servers": ["pre-save-idp"]})) + + client = MagicMock() + client.get = slow_get + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + try: + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=client, + ): + async with discoverable_endpoints._oauth_metadata_fetch_slot(cache_key): + shared_lock: Final = discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS[cache_key] + waiting: Final = asyncio.create_task(fetch_upstream_oauth_protected_resource(server)) + for _ in range(3): + await asyncio.sleep(0) + assert not started.is_set() and not waiting.done() + invalidate_oauth_metadata_cache(server.server_id) + assert discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.get(cache_key) is shared_lock + assert discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + await started.wait() + invalidate_oauth_metadata_cache(server.server_id) + release.set() + assert await waiting == {"authorization_servers": ["pre-save-idp"]} + assert cache_key not in discoverable_endpoints._OAUTH_METADATA_CACHE + assert not discoverable_endpoints._oauth_metadata_fetch_in_flight(server.server_id) + finally: + discoverable_endpoints._OAUTH_METADATA_CACHE.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_FETCHERS.pop(cache_key, None) + discoverable_endpoints._OAUTH_METADATA_GENERATIONS.pop(server.server_id, None) + + def test_invalidating_an_idle_server_leaves_no_generation_behind(): from litellm.proxy._experimental.mcp_server import discoverable_endpoints from litellm.proxy._experimental.mcp_server.discoverable_endpoints import invalidate_oauth_metadata_cache From a3c3ca08cd43ba03d4864c259bab0468e64099e3 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 23:04:40 +0000 Subject: [PATCH 6/7] fix(mcp): keep a held OAuth metadata lock registered even when no fetcher slot claims it Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp_server/discoverable_endpoints.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index dd72328732b..7a0f59c3c2b 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -142,7 +142,7 @@ def _prune_oauth_metadata_cache(now: float | None = None) -> None: # Drop locks whose cache entry has been evicted and that nobody holds or # waits on; the rest stay so in-flight callers continue to coalesce. for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS): - if cache_key in _OAUTH_METADATA_CACHE or cache_key in _OAUTH_METADATA_FETCHERS: + if cache_key in _OAUTH_METADATA_CACHE or not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) @@ -154,6 +154,13 @@ def _oauth_metadata_fetch_in_flight(server_id: str) -> bool: return any(cache_key[0] == server_id for cache_key in _OAUTH_METADATA_FETCHERS) +def _oauth_metadata_lock_idle(cache_key: tuple[str, str]) -> bool: + if cache_key in _OAUTH_METADATA_FETCHERS: + return False + lock: Final = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key) + return lock is None or not lock.locked() + + @asynccontextmanager async def _oauth_metadata_fetch_slot(cache_key: tuple[str, str]) -> AsyncIterator[None]: _OAUTH_METADATA_FETCHERS[cache_key] = _OAUTH_METADATA_FETCHERS.get(cache_key, 0) + 1 @@ -177,7 +184,7 @@ def invalidate_oauth_metadata_cache(server_id: str) -> None: for cache_key in [key for key in _OAUTH_METADATA_CACHE if key[0] == server_id]: del _OAUTH_METADATA_CACHE[cache_key] for cache_key in [key for key in _OAUTH_METADATA_FETCH_LOCKS if key[0] == server_id]: - if cache_key in _OAUTH_METADATA_FETCHERS: + if not _oauth_metadata_lock_idle(cache_key): continue _OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None) From 83ae2a9bbe3eaa7064f958ba3f9633475d710ea6 Mon Sep 17 00:00:00 2001 From: yucheng Date: Mon, 28 Sep 2026 23:43:18 +0000 Subject: [PATCH 7/7] test(mcp): prove a peer worker drops stale upstream OAuth metadata after a save elsewhere Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../mcp/test_oauth_configuration.py | 24 +++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/tests/integration/mcp/test_oauth_configuration.py b/tests/integration/mcp/test_oauth_configuration.py index 3a023f6d38e..fe2b1069f04 100644 --- a/tests/integration/mcp/test_oauth_configuration.py +++ b/tests/integration/mcp/test_oauth_configuration.py @@ -166,6 +166,14 @@ def _advertised_authorization_servers(gateway: Gateway, alias: str) -> tuple[str return tuple(TypeAdapter(list[str]).validate_python(response.json()["authorization_servers"])) +def _eventually_advertises(gateway: Gateway, alias: str, issuer: str) -> None: + eventually( + lambda: gateway.client.get(f"/.well-known/oauth-protected-resource/{alias}/mcp"), + lambda response: response.status_code == 200 and response.json()["authorization_servers"] == [issuer], + seconds=40, + ) + + def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gateway: Gateway) -> None: moved: Final = threading.Event() with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: @@ -183,6 +191,22 @@ def test_saving_a_pass_through_server_refetches_its_upstream_oauth_metadata(gate ) +def test_peer_worker_stops_advertising_the_old_idp_after_a_save_on_another_worker( + gateway: Gateway, peer: Gateway +) -> None: + moved: Final = threading.Event() + with wire_server(_idp_upstream(lambda: wire.url, moved)) as wire, gateway.scenario() as scenario: + alias: Final = "pt" + uuid.uuid4().hex[:8] + identity: Final = _register_pass_through(scenario, wire, alias) + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-before",) + _eventually_advertises(peer, alias, wire.url + "/idp-before") + moved.set() + saved: Final = gateway.request("PUT", "/v1/mcp/server", {"server_id": identity, "description": "IdP moved"}) + assert saved.status_code == 202, saved.text + assert _advertised_authorization_servers(gateway, alias) == (wire.url + "/idp-after",) + _eventually_advertises(peer, alias, wire.url + "/idp-after") + + def test_metadata_fetched_before_a_save_cannot_repopulate_the_cache_after_it(gateway: Gateway) -> None: moved: Final = threading.Event() hold: Final = _Hold()