diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 49cd8cf6fb9..b8099c97861 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -3801,7 +3801,7 @@ class MCPServerManager: if server is None: verbose_logger.warning("MCP Server %s not found", server_id) return [] - return await self._get_tools_from_server(server) + return list(await self._get_tools_from_server(server)) except Exception as e: verbose_logger.warning("Failed to get tools from server %s: %s", server_id, e) return [] @@ -3838,11 +3838,13 @@ class MCPServerManager: server_auth_header: Final = _server_auth_header_for(server, mcp_server_auth_headers, mcp_auth_header) try: - tools: Final = await self._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - user_api_key_auth=user_api_key_auth, - record_listing=True, + tools: Final = list( + await self._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + user_api_key_auth=user_api_key_auth, + record_listing=True, + ) ) return tools except Exception as e: @@ -4350,12 +4352,11 @@ class MCPServerManager: # applied (e.g. "test_petstore-getinventory"). Do NOT pass them # through _create_prefixed_tools — that would add the prefix a second # time producing "test_petstore-test_petstore-getinventory". - unprefixed_tools: Final = guarded_openapi - self._record_listed_tools( - server, unprefixed_tools, listed_caller, listed_generation, record_listing=record_listing + self.record_listed_tools( + server, guarded_openapi, listed_caller, listed_generation, record_listing=record_listing ) if not add_prefix: - return unprefixed_tools + return guarded_openapi return [t.model_copy(update={"name": registered_names[t.name]}) for t in guarded_openapi] else: tools = await self._fetch_tools_with_timeout(client, server.name) @@ -4371,7 +4372,7 @@ class MCPServerManager: prefixed_or_original_tools: Final = self._create_prefixed_tools( guarded_tools, server, add_prefix=add_prefix ) - self._record_listed_tools( + self.record_listed_tools( server, guarded_tools, listed_caller, listed_generation, record_listing=record_listing ) @@ -4478,17 +4479,6 @@ class MCPServerManager: return self._listed_tools_generations.get(server_id, 0) def record_listed_tools( - self, - server: MCPServer, - tools: Sequence[MCPTool], - caller: ListedToolsCaller | None, - generation: int, - *, - record_listing: bool = True, - ) -> None: - self._record_listed_tools(server, tools, caller, generation, record_listing=record_listing) - - def _record_listed_tools( self, server: MCPServer, tools: Sequence[MCPTool], diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 1dea44a84f0..42cca179cd8 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -1125,18 +1125,20 @@ async def _get_tools_from_mcp_servers( from litellm.proxy.proxy_server import proxy_logging_obj listed_generation: Final = global_mcp_server_manager.listed_tools_generation(server.server_id) - tools: Final = await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - oauth2_headers=oauth2_headers, - proxy_logging_obj=proxy_logging_obj, - catalog_auth_header=catalog_auth_header, - record_listing=False, + tools: Final = list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + oauth2_headers=oauth2_headers, + proxy_logging_obj=proxy_logging_obj, + catalog_auth_header=catalog_auth_header, + record_listing=False, + ) ) filtered_tools = filter_tools_by_allowed_tools(tools, server) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 845dcaa1f16..534eda07292 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -705,16 +705,18 @@ if MCP_AVAILABLE: *, record_listing: bool, ) -> list[MCPTool]: - return await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=False, - raw_headers=raw_headers, - client_ip=client_ip, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - record_listing=record_listing, + return list( + await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=False, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + record_listing=record_listing, + ) ) async def _get_tools_for_single_server( diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py index cb017afbea5..e7602d7f825 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -7197,7 +7197,7 @@ class TestMCPServerManager: manager.tool_name_to_mcp_server_name_mapping["test_tool"] = "test-server" manager.tool_name_to_mcp_server_name_mapping["test-server-test_tool"] = "test-server" manager._create_prefixed_tools(listed_tools, server) - manager._record_listed_tools(server, listed_tools, caller) + manager.record_listed_tools(server, listed_tools, caller) mock_client = AsyncMock() mock_client.call_tool.return_value = MagicMock(spec=CallToolResult, content=[], isError=False) @@ -7281,8 +7281,8 @@ class TestMCPServerManager: def test_get_listed_tool_resolves_the_bare_name_from_the_latest_listing(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) - manager._record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v1", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="v2", inputSchema={})], None) latest = manager.get_listed_tool(server, "echo") assert latest is not None and latest.description == "v2" @@ -7294,7 +7294,7 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv-id", name="srv", alias="srv", transport=MCPTransport.http, url="http://srv") caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-user", user_id="alice")) - manager._record_listed_tools( + manager.record_listed_tools( server, [ MCPTool(name="foo", description="Fetches foo records", inputSchema={"type": "object"}), @@ -7355,8 +7355,8 @@ class TestMCPServerManager: manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") other = MCPServer(server_id="other", name="other", transport=MCPTransport.http, url="http://other") - manager._record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) - manager._record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="old", inputSchema={})], None) + manager.record_listed_tools(other, [MCPTool(name="ping", description="kept", inputSchema={})], None) manager._invalidate_server_definition_caches(server.server_id) @@ -7418,7 +7418,7 @@ class TestMCPServerManager: caller = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm", user_id="lister")) async def register_while_a_listing_records(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7463,7 +7463,7 @@ class TestMCPServerManager: return [Prompt(name="greet")] async def register_while_discovery_fills(server: MCPServer, *, initialize_mapping: bool = True) -> None: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="search", description="pre-save", inputSchema={})], caller, @@ -7497,7 +7497,7 @@ class TestMCPServerManager: async def test_user_oauth_refresh_keeps_listed_tools(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="echo", description="shared", inputSchema={})], None) await manager.invalidate_user_oauth_token_cache("alice", server.server_id) @@ -7517,12 +7517,12 @@ class TestMCPServerManager: bob = UserAPIKeyAuth(user_id="bob", token="hashed-bob") alice_schema = {"type": "object", "properties": {"path": {"type": "string"}}} bob_schema = {"type": "object", "properties": {"path": {"type": "string"}, "site": {"type": "string"}}} - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="alice view", inputSchema=alice_schema)], ListedToolsCaller(user_api_key_auth=alice), ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description="bob view", inputSchema=bob_schema)], ListedToolsCaller(user_api_key_auth=bob), @@ -7539,7 +7539,7 @@ class TestMCPServerManager: assert manager.get_listed_tool(server, "read", carol) is None shared = MCPServer(server_id="shared", name="shared", transport=MCPTransport.http, url="http://shared") - manager._record_listed_tools( + manager.record_listed_tools( shared, [MCPTool(name="echo", description="everyone", inputSchema={})], ListedToolsCaller(user_api_key_auth=alice), @@ -7595,8 +7595,8 @@ class TestMCPServerManager: server = MCPServer( **{"server_id": "srv", "name": "srv", "transport": MCPTransport.http, "url": "http://srv", **server_kwargs} ) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) - manager._record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog A", inputSchema={})], caller_a) + manager.record_listed_tools(server, [MCPTool(name="turn", description="Catalog B", inputSchema={})], caller_b) for_a = manager.get_listed_tool(server, "turn", caller_a) for_b = manager.get_listed_tool(server, "turn", caller_b) @@ -7607,7 +7607,7 @@ class TestMCPServerManager: def test_shared_server_ignores_headers_it_never_forwards(self): manager = MCPServerManager() server = MCPServer(server_id="srv", name="srv", transport=MCPTransport.http, url="http://srv") - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="everyone", inputSchema={})], ListedToolsCaller(raw_headers={"authorization": "Bearer sk-litellm", "x-workspace": "A"}), @@ -7840,7 +7840,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=signer, ): - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="turn", description="alice view", inputSchema={})], alice ) assert manager.get_listed_tool(server, "turn", bob) is None @@ -7860,7 +7860,7 @@ class TestMCPServerManager: "litellm.proxy.guardrails.guardrail_hooks.mcp_jwt_signer.mcp_jwt_signer.get_mcp_jwt_signer", return_value=MagicMock(), ): - manager._record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="turn", description="slot a", inputSchema={})], alice) assert manager.get_listed_tool(server, "turn", bob) is None same_key = ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id="same-user", api_key="sk-alpha")) @@ -7878,7 +7878,7 @@ class TestMCPServerManager: team_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], team_one ) @@ -7896,7 +7896,7 @@ class TestMCPServerManager: alice_in_two: Final = ListedToolsCaller( user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id="alice", team_id="team-two") ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="foo", description="Fetch rows FLAGWORD", inputSchema={})], alice_in_one ) @@ -7917,7 +7917,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key=None, user_id=None, team_id="team-one"), raw_headers={"authorization": "Bearer jwt-bob"}, ) - manager._record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) + manager.record_listed_tools(server, [MCPTool(name="foo", description="alice view", inputSchema={})], alice) assert manager.get_listed_tool(server, "foo", bob) is None listed: Final = manager.get_listed_tool(server, "foo", alice) @@ -7976,7 +7976,7 @@ class TestMCPServerManager: user_api_key_auth=UserAPIKeyAuth(api_key="sk-master"), raw_headers={"x-litellm-api-key": "Bearer sk-master", "authorization": "Bearer UP-B"}, ) - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="lookup", description="Workspace A lookup FLAGWORD", inputSchema={})], caller_a ) @@ -8053,18 +8053,18 @@ class TestMCPServerManager: url="http://srv", auth_type=MCPAuth.oauth2_token_exchange, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) + manager.record_listed_tools(server, [MCPTool(name="read", description="shared", inputSchema={})], None) callers = [ ListedToolsCaller(user_api_key_auth=UserAPIKeyAuth(user_id=f"u{i}", api_key=f"k{i}")) for i in range(_LISTED_TOOLS_CALLERS_PER_SERVER + 1) ] for caller in callers: - manager._record_listed_tools( + manager.record_listed_tools( server, [MCPTool(name="read", description=caller.user_api_key_auth.user_id, inputSchema={})], caller, ) - manager._record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) + manager.record_listed_tools(server, [MCPTool(name="read", description="u1 again", inputSchema={})], callers[1]) assert manager.get_listed_tool(server, "read", callers[0]) is None second = manager.get_listed_tool(server, "read", callers[1]) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 556fe0b2536..3b53d023af3 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -8115,7 +8115,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_listed_entry_and_nothing patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), ): never_listed_tool, never_listed_data = await call() - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="list_pets", description="ADMIN DESC", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8156,7 +8156,7 @@ async def test_execute_mcp_tool_hands_openapi_hooks_the_guarded_catalog_entry_cl ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=pinned_schema)], ListedToolsCaller(user_api_key_auth=alice), @@ -8206,12 +8206,12 @@ async def test_execute_mcp_tool_hands_openapi_hooks_each_callers_own_listed_entr manager = mcp_module.global_mcp_server_manager guarded = UserAPIKeyAuth(api_key="sk-guarded", user_id="alice") opted_out = UserAPIKeyAuth(api_key="sk-opted-out", user_id="bob") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a [MASKED] pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=guarded), ) - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="getpetbyid", description="Find a SECRET pet", inputSchema=schema)], ListedToolsCaller(user_api_key_auth=opted_out), @@ -8305,7 +8305,7 @@ async def test_execute_mcp_tool_hands_hooks_nothing_for_a_never_listed_operation ) manager = mcp_module.global_mcp_server_manager alice = UserAPIKeyAuth(api_key="sk-user", user_id="alice") - manager._record_listed_tools( + manager.record_listed_tools( petstore, [MCPTool(name="get_pet", description="Fetches pet records. FLAGWORD", inputSchema={"type": "object"})], ListedToolsCaller(user_api_key_auth=alice),