feat: add ability to remove tools (#1322)
Co-authored-by: David Soria Parra <167242713+dsp-ant@users.noreply.github.com> Co-authored-by: Felix Weinberger <3823880+felixweinberger@users.noreply.github.com> Co-authored-by: Max Isbey <224885523+maxisbey@users.noreply.github.com> Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -391,6 +391,17 @@ class FastMCP(Generic[LifespanResultT]):
|
||||
structured_output=structured_output,
|
||||
)
|
||||
|
||||
def remove_tool(self, name: str) -> None:
|
||||
"""Remove a tool from the server by name.
|
||||
|
||||
Args:
|
||||
name: The name of the tool to remove
|
||||
|
||||
Raises:
|
||||
ToolError: If the tool does not exist
|
||||
"""
|
||||
self._tool_manager.remove_tool(name)
|
||||
|
||||
def tool(
|
||||
self,
|
||||
name: str | None = None,
|
||||
|
||||
@@ -70,6 +70,12 @@ class ToolManager:
|
||||
self._tools[tool.name] = tool
|
||||
return tool
|
||||
|
||||
def remove_tool(self, name: str) -> None:
|
||||
"""Remove a tool by name."""
|
||||
if name not in self._tools:
|
||||
raise ToolError(f"Unknown tool: {name}")
|
||||
del self._tools[name]
|
||||
|
||||
async def call_tool(
|
||||
self,
|
||||
name: str,
|
||||
|
||||
@@ -603,6 +603,80 @@ class TestServerTools:
|
||||
assert result.isError is False
|
||||
assert result.structuredContent == {"theme": "dark", "language": "en", "timezone": "UTC"}
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remove_tool(self):
|
||||
"""Test removing a tool from the server."""
|
||||
mcp = FastMCP()
|
||||
mcp.add_tool(tool_fn)
|
||||
|
||||
# Verify tool exists
|
||||
assert len(mcp._tool_manager.list_tools()) == 1
|
||||
|
||||
# Remove the tool
|
||||
mcp.remove_tool("tool_fn")
|
||||
|
||||
# Verify tool is removed
|
||||
assert len(mcp._tool_manager.list_tools()) == 0
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remove_nonexistent_tool(self):
|
||||
"""Test that removing a non-existent tool raises ToolError."""
|
||||
from mcp.server.fastmcp.exceptions import ToolError
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
with pytest.raises(ToolError, match="Unknown tool: nonexistent"):
|
||||
mcp.remove_tool("nonexistent")
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remove_tool_and_list(self):
|
||||
"""Test that a removed tool doesn't appear in list_tools."""
|
||||
mcp = FastMCP()
|
||||
mcp.add_tool(tool_fn)
|
||||
mcp.add_tool(error_tool_fn)
|
||||
|
||||
# Verify both tools exist
|
||||
async with client_session(mcp._mcp_server) as client:
|
||||
tools = await client.list_tools()
|
||||
assert len(tools.tools) == 2
|
||||
tool_names = [t.name for t in tools.tools]
|
||||
assert "tool_fn" in tool_names
|
||||
assert "error_tool_fn" in tool_names
|
||||
|
||||
# Remove one tool
|
||||
mcp.remove_tool("tool_fn")
|
||||
|
||||
# Verify only one tool remains
|
||||
async with client_session(mcp._mcp_server) as client:
|
||||
tools = await client.list_tools()
|
||||
assert len(tools.tools) == 1
|
||||
assert tools.tools[0].name == "error_tool_fn"
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_remove_tool_and_call(self):
|
||||
"""Test that calling a removed tool fails appropriately."""
|
||||
mcp = FastMCP()
|
||||
mcp.add_tool(tool_fn)
|
||||
|
||||
# Verify tool works before removal
|
||||
async with client_session(mcp._mcp_server) as client:
|
||||
result = await client.call_tool("tool_fn", {"x": 1, "y": 2})
|
||||
assert not result.isError
|
||||
content = result.content[0]
|
||||
assert isinstance(content, TextContent)
|
||||
assert content.text == "3"
|
||||
|
||||
# Remove the tool
|
||||
mcp.remove_tool("tool_fn")
|
||||
|
||||
# Verify calling removed tool returns an error
|
||||
async with client_session(mcp._mcp_server) as client:
|
||||
result = await client.call_tool("tool_fn", {"x": 1, "y": 2})
|
||||
assert result.isError
|
||||
content = result.content[0]
|
||||
assert isinstance(content, TextContent)
|
||||
assert "Unknown tool" in content.text
|
||||
|
||||
|
||||
class TestServerResources:
|
||||
@pytest.mark.anyio
|
||||
|
||||
@@ -633,3 +633,116 @@ class TestStructuredOutput:
|
||||
# Test converted result
|
||||
result = await manager.call_tool("get_scores", {})
|
||||
assert result == expected
|
||||
|
||||
|
||||
class TestRemoveTools:
|
||||
"""Test tool removal functionality in the tool manager."""
|
||||
|
||||
def test_remove_existing_tool(self):
|
||||
"""Test removing an existing tool."""
|
||||
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Add two numbers."""
|
||||
return a + b
|
||||
|
||||
manager = ToolManager()
|
||||
manager.add_tool(add)
|
||||
|
||||
# Verify tool exists
|
||||
assert manager.get_tool("add") is not None
|
||||
assert len(manager.list_tools()) == 1
|
||||
|
||||
# Remove the tool - should not raise any exception
|
||||
manager.remove_tool("add")
|
||||
|
||||
# Verify tool is removed
|
||||
assert manager.get_tool("add") is None
|
||||
assert len(manager.list_tools()) == 0
|
||||
|
||||
def test_remove_nonexistent_tool(self):
|
||||
"""Test removing a non-existent tool raises ToolError."""
|
||||
manager = ToolManager()
|
||||
|
||||
with pytest.raises(ToolError, match="Unknown tool: nonexistent"):
|
||||
manager.remove_tool("nonexistent")
|
||||
|
||||
def test_remove_tool_from_multiple_tools(self):
|
||||
"""Test removing one tool when multiple tools exist."""
|
||||
|
||||
def add(a: int, b: int) -> int:
|
||||
"""Add two numbers."""
|
||||
return a + b
|
||||
|
||||
def multiply(a: int, b: int) -> int:
|
||||
"""Multiply two numbers."""
|
||||
return a * b
|
||||
|
||||
def divide(a: int, b: int) -> float:
|
||||
"""Divide two numbers."""
|
||||
return a / b
|
||||
|
||||
manager = ToolManager()
|
||||
manager.add_tool(add)
|
||||
manager.add_tool(multiply)
|
||||
manager.add_tool(divide)
|
||||
|
||||
# Verify all tools exist
|
||||
assert len(manager.list_tools()) == 3
|
||||
assert manager.get_tool("add") is not None
|
||||
assert manager.get_tool("multiply") is not None
|
||||
assert manager.get_tool("divide") is not None
|
||||
|
||||
# Remove middle tool
|
||||
manager.remove_tool("multiply")
|
||||
|
||||
# Verify only multiply is removed
|
||||
assert len(manager.list_tools()) == 2
|
||||
assert manager.get_tool("add") is not None
|
||||
assert manager.get_tool("multiply") is None
|
||||
assert manager.get_tool("divide") is not None
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_call_removed_tool_raises_error(self):
|
||||
"""Test that calling a removed tool raises ToolError."""
|
||||
|
||||
def greet(name: str) -> str:
|
||||
"""Greet someone."""
|
||||
return f"Hello, {name}!"
|
||||
|
||||
manager = ToolManager()
|
||||
manager.add_tool(greet)
|
||||
|
||||
# Verify tool works before removal
|
||||
result = await manager.call_tool("greet", {"name": "World"})
|
||||
assert result == "Hello, World!"
|
||||
|
||||
# Remove the tool
|
||||
manager.remove_tool("greet")
|
||||
|
||||
# Verify calling removed tool raises error
|
||||
with pytest.raises(ToolError, match="Unknown tool: greet"):
|
||||
await manager.call_tool("greet", {"name": "World"})
|
||||
|
||||
def test_remove_tool_case_sensitive(self):
|
||||
"""Test that tool removal is case-sensitive."""
|
||||
|
||||
def test_func() -> str:
|
||||
"""Test function."""
|
||||
return "test"
|
||||
|
||||
manager = ToolManager()
|
||||
manager.add_tool(test_func)
|
||||
|
||||
# Verify tool exists
|
||||
assert manager.get_tool("test_func") is not None
|
||||
|
||||
# Try to remove with different case - should raise ToolError
|
||||
with pytest.raises(ToolError, match="Unknown tool: Test_Func"):
|
||||
manager.remove_tool("Test_Func")
|
||||
|
||||
# Verify original tool still exists
|
||||
assert manager.get_tool("test_func") is not None
|
||||
|
||||
# Remove with correct case
|
||||
manager.remove_tool("test_func")
|
||||
assert manager.get_tool("test_func") is None
|
||||
|
||||
Reference in New Issue
Block a user