fastmcp/tests/server/middleware/test_tool_injection.py
Jeremiah Lowin 41ec7ee06d
SEP-1577: Sampling with tools (#2551)
* MCP → SDK (vocab change only)

* WIP: Sampling API with SamplingResult[T] and result_type

* SEP-1577: Sampling with tools

- Add tools and result_type parameters to ctx.sample()
- Update OpenAI handler for tool content types
- Client advertises sampling.tools capability by default
- Collect tool results into single message with list content

* Fix tool result content handling in OpenAI handler

* Remove @sampling_tool decorator - pass functions directly to sample()

Functions passed to ctx.sample(tools=[...]) are now auto-converted
via SamplingTool.from_function(). Users can still use that method
directly for custom name/description overrides.

* Remove auto-conversion of MCP tools to sampling tools

Users want MCP tools passed to ctx.sample() to go through the full MCP
machinery (middleware, native responses) rather than being auto-converted
to direct function calls. Now only SamplingTool and plain callables are
accepted - passing a FastMCP Tool raises a clear TypeError.

Also bumps mcp dependency to >=1.24.0 for required sampling features.

* Refactor sampling API: replace sample_iter() with sample_step()

Replace the mutable SampleRun/sample_iter() pattern with a simpler stateless
sample_step() function. sample_step() makes a single LLM call and returns a
SampleStep with the response and history. sample() now loops sample_step()
internally.

Key changes:
- Add sample_step() for fine-grained control over the sampling loop
- Remove SampleRun class and sample_iter() method
- Structured output uses tool description only (no prompt modification)
- execute_tools parameter controls automatic vs manual tool execution

* Address CodeRabbit nitpicks

* Address CodeRabbit review feedback for sampling tools

- Fix temperature=0.0 being dropped due to falsy evaluation
- Add ToolChoice.name support for forcing specific tools
- Replace assert statements with explicit RuntimeError checks
- Add mask_error_details parameter to sample()/sample_step() with ToolError escape hatch
- Fix hasattr patterns with proper isinstance checks
- Document mask_error_details and add OpenAI prerequisites to docs

* Address additional CodeRabbit review feedback

- Catch ValidationError specifically instead of bare Exception
- Update result_type docs to mention dataclasses and basic types
- Raise ValueError for unknown tool_choice modes
- Validate sampling_handler_behavior to catch typos
- Remove ToolChoice.name handling (not part of MCP spec)
- Validate tool_choice string in sample_step()

* Review fixes for sampling tools PR

- Remove internal functions from sampling __init__.py exports
- Remove fragile is_text property, use not is_tool_use instead
- Inline call_client into context.py, remove from run.py
- Fix SamplingMessage docs to use TextContent
- Handle result.text being None in doc examples
- Simplify client sampling docs to recommend OpenAISamplingHandler
- Add sampling_capabilities override documentation
- Raise iteration limit from 50 to 100
- Remove _parse_model_preferences duplication
- Use AsyncOpenAI in OpenAISamplingHandler
- Fix tool_choice docstring

* Fix OpenAI handler tests to use AsyncOpenAI

* Address remaining CodeRabbit review comments

- Fix message ordering in OpenAI handler: tool results now correctly
  follow assistant message with tool_calls
- sample_step() now always includes assistant message in history
- Raise ValueError on JSON parse errors instead of silent {}
- Add has_sampling capability check when behavior is None
- Raise RuntimeError when structured output receives text response
- Wrap primitive result_type schemas in object wrapper
- Fix docs example using invalid SamplingMessage construction
- Add comprehensive client_sampling_test.py example

* Add return type annotation to OpenAISamplingHandler.__init__

* Use explicit 'is not None' check for sampling_capabilities defaulting

---------

Co-authored-by: claude[bot] <41898282+claude[bot]@users.noreply.github.com>
Co-authored-by: Bill Easton <strawgate@users.noreply.github.com>
2025-12-14 13:51:05 -05:00

498 lines
18 KiB
Python

"""Tests for tool injection middleware."""
import math
import pytest
from inline_snapshot import snapshot
from mcp.types import TextContent
from mcp.types import Tool as SDKTool
from fastmcp import FastMCP
from fastmcp.client import Client
from fastmcp.client.client import CallToolResult
from fastmcp.client.transports import FastMCPTransport
from fastmcp.server.middleware.tool_injection import (
PromptToolMiddleware,
ResourceToolMiddleware,
ToolInjectionMiddleware,
)
from fastmcp.tools.tool import FunctionTool, Tool
def multiply_fn(a: int, b: int) -> int:
"""Multiply two numbers."""
return a * b
def divide_fn(a: int, b: int) -> float:
"""Divide two numbers."""
if b == 0:
raise ValueError("Cannot divide by zero")
return a / b
multiply_tool = Tool.from_function(fn=multiply_fn, name="multiply", tags={"math"})
divide_tool = Tool.from_function(fn=divide_fn, name="divide", tags={"math"})
class TestToolInjectionMiddleware:
"""Tests with real FastMCP server."""
@pytest.fixture
def base_server(self):
"""Create a base FastMCP server."""
mcp = FastMCP("BaseServer")
@mcp.tool
def add(a: int, b: int) -> int:
"""Add two numbers."""
return a + b
@mcp.tool
def subtract(a: int, b: int) -> int:
"""Subtract two numbers."""
return a - b
return mcp
async def test_list_tools_includes_injected_tools(self, base_server: FastMCP):
"""Test that list_tools returns both base and injected tools."""
injected_tools: list[FunctionTool] = [
multiply_tool,
divide_tool,
]
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=injected_tools
)
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
tools: list[SDKTool] = await client.list_tools()
# Should have all tools: multiply, divide, add, subtract
assert len(tools) == 4
tool_names: list[str] = [tool.name for tool in tools]
assert "multiply" in tool_names
assert "divide" in tool_names
assert "add" in tool_names
assert "subtract" in tool_names
async def test_call_injected_tool(self, base_server: FastMCP):
"""Test that injected tools can be called successfully."""
injected_tools: list[FunctionTool] = [multiply_tool]
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=injected_tools
)
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
result: CallToolResult = await client.call_tool(
name="multiply", arguments={"a": 7, "b": 6}
)
assert result.structured_content is not None
assert result.structured_content["result"] == 42 # type: ignore[attr-defined]
async def test_call_base_tool_still_works(self, base_server: FastMCP):
"""Test that base server tools still work after injecting tools."""
injected_tools: list[FunctionTool] = [multiply_tool]
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=injected_tools
)
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
result: CallToolResult = await client.call_tool(
name="add", arguments={"a": 10, "b": 5}
)
assert result.structured_content is not None
assert result.structured_content["result"] == 15 # type: ignore[attr-defined]
async def test_injected_tool_error_handling(self, base_server: FastMCP):
"""Test that errors in injected tools are properly handled."""
injected_tools: list[FunctionTool] = [divide_tool]
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=injected_tools
)
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
with pytest.raises(Exception, match="Cannot divide by zero"):
_ = await client.call_tool(name="divide", arguments={"a": 10, "b": 0})
async def test_multiple_tool_injections(self, base_server: FastMCP):
"""Test multiple tool injection middlewares can be stacked."""
def power(a: int, b: int) -> int:
"""Raise a to the power of b."""
return int(math.pow(float(a), float(b)))
def modulo(a: int, b: int) -> int:
"""Calculate a modulo b."""
return a % b
middleware1 = ToolInjectionMiddleware(
tools=[Tool.from_function(fn=power, name="power")]
)
middleware2 = ToolInjectionMiddleware(
tools=[Tool.from_function(fn=modulo, name="modulo")]
)
base_server.add_middleware(middleware1)
base_server.add_middleware(middleware2)
async with Client(base_server) as client:
tools = await client.list_tools()
# Should have all tools
assert len(tools) == 4
tool_names = [tool.name for tool in tools]
assert "power" in tool_names
assert "modulo" in tool_names
assert "add" in tool_names
assert "subtract" in tool_names
# Test that both injected tools work
async with Client(base_server) as client:
power_result = await client.call_tool("power", {"a": 2, "b": 3})
assert power_result.structured_content is not None
assert power_result.structured_content["result"] == 8 # type: ignore[attr-defined]
modulo_result = await client.call_tool("modulo", {"a": 10, "b": 3})
assert modulo_result.structured_content is not None
assert modulo_result.structured_content["result"] == 1 # type: ignore[attr-defined]
async def test_injected_tool_with_complex_return_type(self, base_server: FastMCP):
"""Test injected tools with complex return types."""
def calculate_stats(numbers: list[int]) -> dict[str, int | float]:
"""Calculate statistics for a list of numbers."""
return {
"sum": sum(numbers),
"average": sum(numbers) / len(numbers),
"min": min(numbers),
"max": max(numbers),
"count": len(numbers),
}
middleware = ToolInjectionMiddleware(
tools=[Tool.from_function(fn=calculate_stats, name="calculate_stats")]
)
base_server.add_middleware(middleware)
async with Client(base_server) as client:
result = await client.call_tool(
"calculate_stats", {"numbers": [1, 2, 3, 4, 5]}
)
assert result.structured_content is not None
assert isinstance(result.structured_content, dict)
assert result.structured_content == snapshot(
{"sum": 15, "average": 3.0, "min": 1, "max": 5, "count": 5}
)
async def test_injected_tool_metadata_preserved(self, base_server: FastMCP):
"""Test that injected tool metadata is preserved."""
def multiply(a: int, b: int) -> int:
"""Multiply two numbers."""
return a * b
injected_tools = [Tool.from_function(fn=multiply, name="multiply")]
middleware = ToolInjectionMiddleware(tools=injected_tools)
base_server.add_middleware(middleware)
async with Client(base_server) as client:
tools = await client.list_tools()
multiply_tool = next(t for t in tools if t.name == "multiply")
assert multiply_tool.description == "Multiply two numbers."
assert "a" in multiply_tool.inputSchema["properties"]
assert "b" in multiply_tool.inputSchema["properties"]
async def test_injected_tool_does_not_conflict_with_base_tool(
self, base_server: FastMCP
):
"""Test that injected tools with same name as base tools are called correctly."""
def add(a: int, b: int) -> int:
"""Injected add that multiplies instead."""
return a * b
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=[Tool.from_function(fn=add, name="add")]
)
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
result: CallToolResult = await client.call_tool(
name="add", arguments={"a": 5, "b": 3}
)
# Should use the injected tool (multiply behavior)
assert result.structured_content is not None
assert result.structured_content["result"] == 15
async def test_injected_tool_bypass_filtering(self, base_server: FastMCP):
"""Test that injected tools bypass filtering."""
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(
tools=[multiply_tool]
)
base_server.add_middleware(middleware)
base_server.exclude_tags = {"math"}
async with Client[FastMCPTransport](base_server) as client:
tools: list[SDKTool] = await client.list_tools()
tool_names: list[str] = [tool.name for tool in tools]
assert "multiply" in tool_names
async def test_empty_tool_injection(self, base_server: FastMCP):
"""Test that middleware with no tools doesn't affect behavior."""
middleware: ToolInjectionMiddleware = ToolInjectionMiddleware(tools=[])
base_server.add_middleware(middleware)
async with Client[FastMCPTransport](base_server) as client:
tools: list[SDKTool] = await client.list_tools()
result: CallToolResult = await client.call_tool(
name="add", arguments={"a": 3, "b": 4}
)
# Should only have the base tools
assert len(tools) == 2
tool_names: list[str] = [tool.name for tool in tools]
assert "add" in tool_names
assert "subtract" in tool_names
assert result.structured_content is not None
assert result.structured_content["result"] == 7 # type: ignore[attr-defined]
class TestPromptToolMiddleware:
"""Tests for PromptToolMiddleware."""
@pytest.fixture
def server_with_prompts(self):
"""Create a FastMCP server with prompts."""
mcp = FastMCP("PromptServer")
@mcp.tool
def add(a: int, b: int) -> int:
"""Add two numbers."""
return a + b
@mcp.prompt
def greeting(name: str) -> str:
"""Generate a greeting message."""
return f"Hello, {name}!"
@mcp.prompt
def farewell(name: str) -> str:
"""Generate a farewell message."""
return f"Goodbye, {name}!"
return mcp
async def test_prompt_tools_added_to_list(self, server_with_prompts: FastMCP):
"""Test that prompt tools are added to the tool list."""
middleware = PromptToolMiddleware()
server_with_prompts.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_prompts) as client:
tools: list[SDKTool] = await client.list_tools()
tool_names: list[str] = [tool.name for tool in tools]
# Should have: add, list_prompts, get_prompt
assert len(tools) == 3
assert "add" in tool_names
assert "list_prompts" in tool_names
assert "get_prompt" in tool_names
async def test_list_prompts_tool_works(self, server_with_prompts: FastMCP):
"""Test that the list_prompts tool can be called."""
middleware = PromptToolMiddleware()
server_with_prompts.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_prompts) as client:
result: CallToolResult = await client.call_tool(
name="list_prompts", arguments={}
)
assert result.content == snapshot(
[
TextContent(
type="text",
text='[{"name":"greeting","title":null,"description":"Generate a greeting message.","arguments":[{"name":"name","description":null,"required":true}],"icons":null,"_meta":{"_fastmcp":{"tags":[]}}},{"name":"farewell","title":null,"description":"Generate a farewell message.","arguments":[{"name":"name","description":null,"required":true}],"icons":null,"_meta":{"_fastmcp":{"tags":[]}}}]',
)
]
)
assert result.structured_content is not None
assert result.structured_content["result"] == snapshot(
[
{
"name": "greeting",
"title": None,
"description": "Generate a greeting message.",
"arguments": [
{"name": "name", "description": None, "required": True}
],
"icons": None,
"_meta": {"_fastmcp": {"tags": []}},
},
{
"name": "farewell",
"title": None,
"description": "Generate a farewell message.",
"arguments": [
{"name": "name", "description": None, "required": True}
],
"icons": None,
"_meta": {"_fastmcp": {"tags": []}},
},
]
)
async def test_get_prompt_tool_works(self, server_with_prompts: FastMCP):
"""Test that the get_prompt tool can be called."""
middleware = PromptToolMiddleware()
server_with_prompts.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_prompts) as client:
result: CallToolResult = await client.call_tool(
name="get_prompt",
arguments={"name": "greeting", "arguments": {"name": "World"}},
)
# The tool returns the prompt result with structured_content
assert result.content == snapshot(
[
TextContent(
type="text",
text='{"_meta":null,"description":"Generate a greeting message.","messages":[{"role":"user","content":{"type":"text","text":"Hello, World!","annotations":null,"_meta":null}}]}',
)
]
)
assert result.structured_content is not None
assert result.structured_content == snapshot(
{
"_meta": None,
"description": "Generate a greeting message.",
"messages": [
{
"role": "user",
"content": {
"type": "text",
"text": "Hello, World!",
"annotations": None,
"_meta": None,
},
}
],
}
)
class TestResourceToolMiddleware:
"""Tests for ResourceToolMiddleware."""
@pytest.fixture
def server_with_resources(self):
"""Create a FastMCP server with resources."""
mcp = FastMCP("ResourceServer")
@mcp.tool
def add(a: int, b: int) -> int:
"""Add two numbers."""
return a + b
@mcp.resource("file://config.txt")
def config_resource() -> str:
"""Get configuration."""
return "debug=true"
@mcp.resource("file://data.json")
def data_resource() -> str:
"""Get data."""
return '{"count": 42}'
return mcp
async def test_resource_tools_added_to_list(self, server_with_resources: FastMCP):
"""Test that resource tools are added to the tool list."""
middleware = ResourceToolMiddleware()
server_with_resources.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_resources) as client:
tools: list[SDKTool] = await client.list_tools()
tool_names: list[str] = [tool.name for tool in tools]
# Should have: add, list_resources, read_resource
assert len(tools) == 3
assert "add" in tool_names
assert "list_resources" in tool_names
assert "read_resource" in tool_names
async def test_list_resources_tool_works(self, server_with_resources: FastMCP):
"""Test that the list_resources tool can be called."""
middleware = ResourceToolMiddleware()
server_with_resources.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_resources) as client:
result: CallToolResult = await client.call_tool(
name="list_resources", arguments={}
)
assert result.structured_content is not None
assert result.structured_content["result"] == snapshot(
[
{
"name": "config_resource",
"title": None,
"uri": "file://config.txt/",
"description": "Get configuration.",
"mimeType": "text/plain",
"size": None,
"icons": None,
"annotations": None,
"_meta": {"_fastmcp": {"tags": []}},
},
{
"name": "data_resource",
"title": None,
"uri": "file://data.json/",
"description": "Get data.",
"mimeType": "text/plain",
"size": None,
"icons": None,
"annotations": None,
"_meta": {"_fastmcp": {"tags": []}},
},
]
)
async def test_read_resource_tool_works(self, server_with_resources: FastMCP):
"""Test that the read_resource tool can be called."""
middleware = ResourceToolMiddleware()
server_with_resources.add_middleware(middleware)
async with Client[FastMCPTransport](server_with_resources) as client:
result: CallToolResult = await client.call_tool(
name="read_resource", arguments={"uri": "file://config.txt"}
)
assert result.content == snapshot(
[
TextContent(
type="text",
text='[{"content":"debug=true","mime_type":"text/plain"}]',
)
]
)
assert result.structured_content == snapshot(
{"result": [{"content": "debug=true", "mime_type": "text/plain"}]}
)