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>
This commit is contained in:
yucheng 2026-09-28 22:03:22 +00:00
parent cee44351e3
commit 9ce2803f42
6 changed files with 228 additions and 78 deletions

View file

@ -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

View file

@ -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),

View file

@ -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]

View file

@ -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(

View file

@ -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)

View file

@ -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