mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
d620d1ffea
commit
95b8e2aaa8
2 changed files with 79 additions and 17 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue