fix(mcp): scope listed-tool metadata per caller on per-user MCP servers

Servers whose upstream catalog depends on the caller (oauth2 per-user, token exchange, id_jag, per-user env vars, delegated auth) now keep one listed-tool mapping per (user_id, api_key) hash under the server id, so one caller's tools/list cannot supply another caller's description or inputSchema to pre-call guardrails. Shared servers and OpenAPI-backed servers keep a single server-wide entry

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-12 19:32:03 +00:00
parent d620d1ffea
commit 95b8e2aaa8
2 changed files with 79 additions and 17 deletions

View file

@ -1947,7 +1947,9 @@ class MCPServerManager:
"gmail_send_email": "zapier_mcp_server",
}
"""
self._listed_tools_by_server_id: dict[str, Mapping[str, MCPTool]] = {} # mutable-ok: refreshed per tools/list
self._listed_tools_by_server_id: dict[
str, Mapping[str | None, Mapping[str, MCPTool]]
] = {} # mutable-ok: refreshed per tools/list
self._upstream_initialize_instructions_by_server_id: dict[str, str] = {}
# Per-server monotonic timestamp of last upstream prefetch attempt (success,
# empty result, or failure). Used to throttle re-probes for servers that do
@ -4446,15 +4448,15 @@ class MCPServerManager:
unprefixed_tools: Final = [ # mutable-ok: returned through the list[MCPTool] listing contract
t.model_copy(update={"name": t.name[len(registry_prefix) :]}) for t in tools
]
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(
{t.name: t for t in unprefixed_tools}
)
self._record_listed_tools(server, unprefixed_tools, user_api_key_auth)
return tools if add_prefix else unprefixed_tools
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
self._remember_upstream_initialize_instructions(server, client)
prefixed_or_original_tools: Final = self._create_prefixed_tools(tools, server, add_prefix=add_prefix)
prefixed_or_original_tools: Final = self._create_prefixed_tools(
tools, server, add_prefix=add_prefix, user_api_key_auth=user_api_key_auth
)
return prefixed_or_original_tools
@ -4502,6 +4504,29 @@ class MCPServerManager:
self._invalidate_discovery_lists(server_id)
self._listed_tools_by_server_id.pop(server_id, None)
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)
)
def _listed_tools_identity(self, server: MCPServer, user_api_key_auth: UserAPIKeyAuth | None) -> str | None:
if server.spec_path or user_api_key_auth is None or not self._discovers_per_caller(server):
return None
material: Final = json.dumps((user_api_key_auth.user_id, user_api_key_auth.api_key), separators=(",", ":"))
return hashlib.sha256(material.encode()).hexdigest()
def _record_listed_tools(
self, server: MCPServer, tools: Sequence[MCPTool], user_api_key_auth: UserAPIKeyAuth | None
) -> None:
identity: Final = self._listed_tools_identity(server, user_api_key_auth)
listing: Final = MappingProxyType({tool.name: tool for tool in tools})
self._listed_tools_by_server_id[server.server_id] = MappingProxyType(
{**self._listed_tools_by_server_id.get(server.server_id, {}), identity: listing}
)
def _discovery_key(
self,
server: MCPServer,
@ -4512,12 +4537,7 @@ 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 = (
@ -5330,7 +5350,13 @@ class MCPServerManager:
"attempts; the 3-character prefix space is too crowded."
)
def _create_prefixed_tools(self, tools: list[MCPTool], server: MCPServer, add_prefix: bool = True) -> list[MCPTool]:
def _create_prefixed_tools(
self,
tools: list[MCPTool],
server: MCPServer,
add_prefix: bool = True,
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> list[MCPTool]:
"""
Create prefixed tools and update tool mapping.
@ -5362,12 +5388,15 @@ class MCPServerManager:
for spelling in iter_known_tool_name_spellings(original_name, server):
self.tool_name_to_mcp_server_name_mapping[spelling] = prefix
self._listed_tools_by_server_id[server.server_id] = MappingProxyType({tool.name: tool for tool in tools})
self._record_listed_tools(server, tools, user_api_key_auth)
verbose_logger.info("Successfully fetched %s tools from server %s", len(prefixed_tools), server.name)
return prefixed_tools
def get_listed_tool(self, server: MCPServer, name: str) -> MCPTool | None:
listed: Final = self._listed_tools_by_server_id.get(server.server_id)
def get_listed_tool(
self, server: MCPServer, name: str, user_api_key_auth: UserAPIKeyAuth | None = None
) -> MCPTool | None:
identity: Final = self._listed_tools_identity(server, user_api_key_auth)
listed: Final = self._listed_tools_by_server_id.get(server.server_id, {}).get(identity)
if not listed:
return None
return listed.get(name) or listed.get(strip_known_server_prefix(name, server))
@ -6349,7 +6378,7 @@ class MCPServerManager:
server=mcp_server,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
tool=self.get_listed_tool(mcp_server, name),
tool=self.get_listed_tool(mcp_server, name, user_api_key_auth),
)
if "arguments" in hook_result:
arguments = hook_result["arguments"]
@ -6365,7 +6394,7 @@ class MCPServerManager:
proxy_logging_obj=proxy_logging_obj,
start_time=start_time,
litellm_logging_obj=litellm_logging_obj,
tool=self.get_listed_tool(mcp_server, name),
tool=self.get_listed_tool(mcp_server, name, user_api_key_auth),
)
tasks.append(during_hook_task)

View file

@ -6632,6 +6632,39 @@ class TestMCPServerManager:
listed = manager.get_listed_tool(server, "echo")
assert listed is not None and listed.description == "shared"
def test_per_caller_server_keeps_listed_tools_per_identity(self):
manager = MCPServerManager()
server = MCPServer(
server_id="srv",
name="srv",
transport=MCPTransport.http,
url="http://srv",
auth_type=MCPAuth.oauth2_token_exchange,
)
alice = UserAPIKeyAuth(user_id="alice", api_key="hashed-alice")
bob = UserAPIKeyAuth(user_id="bob", api_key="hashed-bob")
alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}}
bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}}
manager._create_prefixed_tools(
[MCPTool(name="read", description="alice view", inputSchema=alice_schema)], server, user_api_key_auth=alice
)
manager._create_prefixed_tools(
[MCPTool(name="read", description="bob view", inputSchema=bob_schema)], server, user_api_key_auth=bob
)
alice_tool = manager.get_listed_tool(server, "srv-read", alice)
bob_tool = manager.get_listed_tool(server, "srv-read", bob)
assert alice_tool is not None and (alice_tool.description, alice_tool.inputSchema) == ("alice view", alice_schema)
assert bob_tool is not None and (bob_tool.description, bob_tool.inputSchema) == ("bob view", bob_schema)
assert manager.get_listed_tool(server, "srv-read", UserAPIKeyAuth(user_id="carol", api_key="k")) is None
shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared")
manager._create_prefixed_tools(
[MCPTool(name="echo", description="everyone", inputSchema={})], shared, user_api_key_auth=alice
)
for_bob = manager.get_listed_tool(shared, "echo", bob)
assert for_bob is not None and for_bob.description == "everyone"
@pytest.mark.asyncio
@pytest.mark.parametrize("add_prefix", [True, False])
async def test_openapi_listing_records_listed_tools(self, add_prefix):