From 6fcd333e52c0cd3cb2f93465361e5eb58eb68147 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Tue, 15 Apr 2025 09:52:31 -0400 Subject: [PATCH] Update tool manager param to key --- src/fastmcp/server/server.py | 2 +- src/fastmcp/tools/tool_manager.py | 28 ++++++++++++++-------------- tests/tools/test_tool_manager.py | 4 ++-- 3 files changed, 17 insertions(+), 17 deletions(-) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 723652386..c53e1ee4e 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -186,7 +186,7 @@ class FastMCP(Generic[LifespanResultT]): self._mcp_server.list_resource_templates()(self._mcp_list_resource_templates) def get_tools(self) -> dict[str, Tool]: - """Get all registered tools, keyed by registered name.""" + """Get all registered tools, indexed by registered key.""" return self._tool_manager.get_tools() def list_tools(self) -> list[Tool]: diff --git a/src/fastmcp/tools/tool_manager.py b/src/fastmcp/tools/tool_manager.py index 764804058..9e0f1ac42 100644 --- a/src/fastmcp/tools/tool_manager.py +++ b/src/fastmcp/tools/tool_manager.py @@ -41,7 +41,7 @@ class ToolManager: return self._tools.get(name) def get_tools(self) -> dict[str, Tool]: - """Get all registered tools, keyed by registered name.""" + """Get all registered tools, indexed by registered key.""" return self._tools def list_tools(self) -> list[Tool]: @@ -50,7 +50,7 @@ class ToolManager: def list_mcp_tools(self) -> list[MCPTool]: """List all registered tools in the format expected by the low-level MCP server.""" - return [tool.to_mcp_tool(name=name) for name, tool in self._tools.items()] + return [tool.to_mcp_tool(name=key) for key, tool in self._tools.items()] def add_tool_from_fn( self, @@ -63,34 +63,34 @@ class ToolManager: tool = Tool.from_function(fn, name=name, description=description, tags=tags) return self.add_tool(tool) - def add_tool(self, tool: Tool, name: str | None = None) -> Tool: + def add_tool(self, tool: Tool, key: str | None = None) -> Tool: """Register a tool with the server.""" - name = name or tool.name - existing = self._tools.get(name) + key = key or tool.name + existing = self._tools.get(key) if existing: if self.duplicate_behavior == "warn": - logger.warning(f"Tool already exists: {name}") - self._tools[name] = tool + logger.warning(f"Tool already exists: {key}") + self._tools[key] = tool elif self.duplicate_behavior == "replace": - self._tools[name] = tool + self._tools[key] = tool elif self.duplicate_behavior == "error": - raise ValueError(f"Tool already exists: {name}") + raise ValueError(f"Tool already exists: {key}") elif self.duplicate_behavior == "ignore": return existing else: - self._tools[name] = tool + self._tools[key] = tool return tool async def call_tool( self, - name: str, + key: str, arguments: dict[str, Any], context: Context[ServerSessionT, LifespanContextT] | None = None, ) -> Any: """Call a tool by name with arguments.""" - tool = self.get_tool(name) + tool = self.get_tool(key) if not tool: - raise ToolError(f"Unknown tool: {name}") + raise ToolError(f"Unknown tool: {key}") return await tool.run(arguments, context=context) @@ -110,5 +110,5 @@ class ToolManager: """ for name, tool in tool_manager._tools.items(): prefixed_name = f"{prefix}{name}" if prefix else name - self.add_tool(tool, name=prefixed_name) + self.add_tool(tool, key=prefixed_name) logger.debug(f'Imported tool "{tool.name}" as "{prefixed_name}"') diff --git a/tests/tools/test_tool_manager.py b/tests/tools/test_tool_manager.py index c86e8b0fb..0df04e6fa 100644 --- a/tests/tools/test_tool_manager.py +++ b/tests/tools/test_tool_manager.py @@ -585,7 +585,7 @@ class TestCustomToolNames: tool = Tool.from_function(fn, name="my_tool") manager = ToolManager() # Store it under a different name - manager.add_tool(tool, name="proxy_tool") + manager.add_tool(tool, key="proxy_tool") # The tool is accessible under the storage name stored = manager.get_tool("proxy_tool") assert stored is not None @@ -682,7 +682,7 @@ class TestCustomToolNames: tool = Tool.from_function(fn, name="my_tool") manager = ToolManager() - manager.add_tool(tool, name="proxy_tool") + manager.add_tool(tool, key="proxy_tool") mcp_tools = manager.list_mcp_tools() assert len(mcp_tools) == 1 assert mcp_tools[0].name == "proxy_tool"