diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index d53921994c8..a1200f1784e 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 8693227feaa..4d1490b76a5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -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):