mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
cee44351e3
commit
9ce2803f42
6 changed files with 228 additions and 78 deletions
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue