Merge remote-tracking branch 'origin/litellm_mcp_discovery_cache_fixes' into litellm_mcp_listed_tool_metadata
Some checks are pending
LiteLLM Rust / rust-lint (push) Waiting to run
LiteLLM Rust / rust-test (push) Waiting to run
LiteLLM Rust / rust-wheel (push) Waiting to run
Terraform Provider / gofmt, vet, build, test (push) Waiting to run
Terraform Provider / Provider endpoints vs proxy OpenAPI schema (push) Waiting to run

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
This commit is contained in:
yucheng 2026-09-28 22:39:06 +00:00
commit 4f56462900
6 changed files with 838 additions and 326 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

@ -4607,39 +4607,23 @@ class MCPServerManager:
self._listed_tools_by_server_id.pop(server_id, None)
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 _listed_tools_identity(self, server: MCPServer, caller: ListedToolsCaller | None) -> str | None:
"""Key the listed-tool cache by every request input that can change the upstream catalog.
Forwarded headers, header-driven stdio env, the caller bearer (forwarded as-is or
exchanged as the OBO subject), and the server-specific auth header all reach
upstream, so two callers differing in any of them may be shown different tools. Shared
servers with none of those stay on the shared (``None``) slot. OpenAPI servers list from
the process-wide registry.
exchanged as the OBO subject), the server-specific auth header, and the per-caller JWT
MCPJWTSigner mints for tools/list all reach upstream, so two callers differing in any of
them may be shown different tools. Shared servers with none of those stay on the shared
(``None``) slot. OpenAPI servers list from the process-wide registry.
"""
if server.spec_path or caller is None:
return None
auth: Final = caller.user_api_key_auth
signed_caller: Final = (
f"{auth.user_id}:{auth.api_key}"
if auth is not None and self._signs_caller_identity_upstream(server)
else None
)
forwarded: Final = dict(self._forwarded_header_values(server, caller.raw_headers)) or None
header_env: Final = self._build_stdio_env(server, caller.raw_headers)
stdio_env: Final = None if header_env == self._build_stdio_env(server) else header_env
@ -4649,7 +4633,19 @@ class MCPServerManager:
else None
)
_, digest = self._discovery_key(server, auth, caller.mcp_auth_header, forwarded, stdio_env, caller_bearer)
return digest
if signed_caller is None:
return digest
return hashlib.sha256(f"{digest}:{signed_caller}".encode()).hexdigest()
@staticmethod
def _signs_caller_identity_upstream(server: MCPServer) -> bool:
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 {}))
@staticmethod
def _forwarded_header_values(
@ -4688,11 +4684,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)