mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 05:54:19 +02:00
Add middleware for all current handlers
This commit is contained in:
parent
c183e3a99c
commit
a42c0c40b0
3 changed files with 440 additions and 116 deletions
|
|
@ -148,7 +148,7 @@ class MCPMiddleware:
|
|||
handler = partial(self.on_list_tools, call_next=handler)
|
||||
case "resources/list":
|
||||
handler = partial(self.on_list_resources, call_next=handler)
|
||||
case "resource-templates/list":
|
||||
case "resources/templates/list":
|
||||
handler = partial(self.on_list_resource_templates, call_next=handler)
|
||||
case "prompts/list":
|
||||
handler = partial(self.on_list_prompts, call_next=handler)
|
||||
|
|
|
|||
|
|
@ -352,34 +352,8 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
|
||||
async def get_resources(self) -> dict[str, Resource]:
|
||||
"""Get all registered resources, indexed by registered key."""
|
||||
if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
|
||||
resources: dict[str, Resource] = {}
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
server_resources = await mounted_server.server.get_resources()
|
||||
# Apply prefix to each resource key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for resource in server_resources.values():
|
||||
resource = resource.with_key(
|
||||
add_resource_prefix(
|
||||
resource.key,
|
||||
mounted_server.prefix,
|
||||
self.resource_prefix_format,
|
||||
)
|
||||
)
|
||||
resources[resource.key] = resource
|
||||
else:
|
||||
resources.update(server_resources)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
resources.update(self._resource_manager.get_resources())
|
||||
self._cache.set("resources", resources)
|
||||
return resources
|
||||
resources = await self._list_resources(apply_middleware=False)
|
||||
return {resource.key: resource for resource in resources}
|
||||
|
||||
async def get_resource(self, key: str) -> Resource:
|
||||
resources = await self.get_resources()
|
||||
|
|
@ -389,39 +363,8 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
|
||||
async def get_resource_templates(self) -> dict[str, ResourceTemplate]:
|
||||
"""Get all registered resource templates, indexed by registered key."""
|
||||
if (
|
||||
templates := self._cache.get("resource_templates")
|
||||
) is self._cache.NOT_FOUND:
|
||||
templates: dict[str, ResourceTemplate] = {}
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
server_templates = (
|
||||
await mounted_server.server.get_resource_templates()
|
||||
)
|
||||
# Apply prefix to each template key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for template in server_templates.values():
|
||||
template = template.with_key(
|
||||
add_resource_prefix(
|
||||
template.key,
|
||||
mounted_server.prefix,
|
||||
self.resource_prefix_format,
|
||||
)
|
||||
)
|
||||
templates[template.key] = template
|
||||
else:
|
||||
templates.update(server_templates)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to get resource templates from mounted server "
|
||||
f"'{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
templates.update(self._resource_manager.get_templates())
|
||||
self._cache.set("resource_templates", templates)
|
||||
return templates
|
||||
templates = await self._list_resource_templates(apply_middleware=False)
|
||||
return {template.key: template for template in templates}
|
||||
|
||||
async def get_resource_template(self, key: str) -> ResourceTemplate:
|
||||
templates = await self.get_resource_templates()
|
||||
|
|
@ -434,30 +377,8 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
List all available prompts.
|
||||
"""
|
||||
|
||||
if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
|
||||
prompts: dict[str, Prompt] = {}
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
server_prompts = await mounted_server.server.get_prompts()
|
||||
# Apply prefix to each prompt key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for prompt in server_prompts.values():
|
||||
prompt = prompt.with_key(
|
||||
f"{mounted_server.prefix}_{prompt.key}"
|
||||
)
|
||||
prompts[prompt.key] = prompt
|
||||
else:
|
||||
prompts.update(server_prompts)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
prompts.update(self._prompt_manager.get_prompts())
|
||||
self._cache.set("prompts", prompts)
|
||||
return prompts
|
||||
prompts = await self._list_prompts(apply_middleware=False)
|
||||
return {prompt.key: prompt for prompt in prompts}
|
||||
|
||||
async def get_prompt(self, key: str) -> Prompt:
|
||||
prompts = await self.get_prompts()
|
||||
|
|
@ -582,21 +503,31 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
return tools
|
||||
|
||||
async def _mcp_list_resources(self) -> list[MCPResource]:
|
||||
logger.debug("Handler called: list_resources")
|
||||
|
||||
with fastmcp.server.context.Context(fastmcp=self):
|
||||
resources = await self._middleware_list_resources()
|
||||
return [
|
||||
resource.to_mcp_resource(uri=resource.key) for resource in resources
|
||||
]
|
||||
|
||||
async def _middleware_list_resources(self) -> list[Resource]:
|
||||
"""
|
||||
List all available resources, in the format expected by the low-level MCP
|
||||
server.
|
||||
|
||||
"""
|
||||
logger.debug("Handler called: list_resources")
|
||||
|
||||
async def _final_handler(
|
||||
async def _handler(
|
||||
context: MiddlewareContext[dict[str, Any]],
|
||||
) -> list[MCPResource]:
|
||||
resources = await self.get_resources()
|
||||
mcp_resources: list[MCPResource] = []
|
||||
for key, resource in resources.items():
|
||||
) -> list[Resource]:
|
||||
resources = await self._list_resources()
|
||||
|
||||
mcp_resources: list[Resource] = []
|
||||
for resource in resources:
|
||||
if self._should_enable_component(resource):
|
||||
mcp_resources.append(resource.to_mcp_resource(uri=key))
|
||||
mcp_resources.append(resource)
|
||||
|
||||
return mcp_resources
|
||||
|
||||
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
|
||||
|
|
@ -610,24 +541,74 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
)
|
||||
|
||||
# Apply the middleware chain.
|
||||
return await self._apply_middleware(mw_context, _final_handler)
|
||||
return await self._apply_middleware(mw_context, _handler)
|
||||
|
||||
async def _list_resources(self, apply_middleware: bool = True) -> list[Resource]:
|
||||
"""
|
||||
List all available resources.
|
||||
"""
|
||||
|
||||
if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
|
||||
resources: list[Resource] = []
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
if apply_middleware:
|
||||
server_resources = (
|
||||
await mounted_server.server._middleware_list_resources()
|
||||
)
|
||||
else:
|
||||
server_resources = await mounted_server.server._list_resources()
|
||||
# Apply prefix to each resource key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for resource in server_resources:
|
||||
resource = resource.with_key(
|
||||
add_resource_prefix(
|
||||
resource.key,
|
||||
mounted_server.prefix,
|
||||
self.resource_prefix_format,
|
||||
)
|
||||
)
|
||||
resources.append(resource)
|
||||
else:
|
||||
resources.extend(server_resources)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
resources.extend(self._resource_manager.get_resources().values())
|
||||
self._cache.set("resources", resources)
|
||||
return resources
|
||||
|
||||
async def _mcp_list_resource_templates(self) -> list[MCPResourceTemplate]:
|
||||
"""
|
||||
List all available resource templates, in the format expected by the low-level
|
||||
MCP server.
|
||||
|
||||
"""
|
||||
logger.debug("Handler called: list_resource_templates")
|
||||
|
||||
async def _final_handler(
|
||||
with fastmcp.server.context.Context(fastmcp=self):
|
||||
templates = await self._middleware_list_resource_templates()
|
||||
return [
|
||||
template.to_mcp_template(uriTemplate=template.key)
|
||||
for template in templates
|
||||
]
|
||||
|
||||
async def _middleware_list_resource_templates(self) -> list[ResourceTemplate]:
|
||||
"""
|
||||
List all available resource templates, in the format expected by the low-level MCP
|
||||
server.
|
||||
|
||||
"""
|
||||
|
||||
async def _handler(
|
||||
context: MiddlewareContext[dict[str, Any]],
|
||||
) -> list[MCPResourceTemplate]:
|
||||
templates = await self.get_resource_templates()
|
||||
mcp_templates: list[MCPResourceTemplate] = []
|
||||
for key, template in templates.items():
|
||||
) -> list[ResourceTemplate]:
|
||||
templates = await self._list_resource_templates()
|
||||
|
||||
mcp_templates: list[ResourceTemplate] = []
|
||||
for template in templates:
|
||||
if self._should_enable_component(template):
|
||||
mcp_templates.append(template.to_mcp_template(uriTemplate=key))
|
||||
mcp_templates.append(template)
|
||||
|
||||
return mcp_templates
|
||||
|
||||
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
|
||||
|
|
@ -636,29 +617,81 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
message={}, # List resource templates doesn't have parameters
|
||||
source="client",
|
||||
type="request",
|
||||
method="resources/list_templates",
|
||||
method="resources/templates/list",
|
||||
fastmcp_context=fastmcp_ctx,
|
||||
)
|
||||
|
||||
# Apply the middleware chain.
|
||||
return await self._apply_middleware(mw_context, _final_handler)
|
||||
return await self._apply_middleware(mw_context, _handler)
|
||||
|
||||
async def _list_resource_templates(
|
||||
self, apply_middleware: bool = True
|
||||
) -> list[ResourceTemplate]:
|
||||
"""
|
||||
List all available resource templates.
|
||||
"""
|
||||
|
||||
if (
|
||||
templates := self._cache.get("resource_templates")
|
||||
) is self._cache.NOT_FOUND:
|
||||
templates: list[ResourceTemplate] = []
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
if apply_middleware:
|
||||
server_templates = await mounted_server.server._middleware_list_resource_templates()
|
||||
else:
|
||||
server_templates = (
|
||||
await mounted_server.server._list_resource_templates()
|
||||
)
|
||||
# Apply prefix to each template key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for template in server_templates:
|
||||
template = template.with_key(
|
||||
add_resource_prefix(
|
||||
template.key,
|
||||
mounted_server.prefix,
|
||||
self.resource_prefix_format,
|
||||
)
|
||||
)
|
||||
templates.append(template)
|
||||
else:
|
||||
templates.extend(server_templates)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"Failed to get resource templates from mounted server "
|
||||
f"'{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
templates.extend(self._resource_manager.get_templates().values())
|
||||
self._cache.set("resource_templates", templates)
|
||||
return templates
|
||||
|
||||
async def _mcp_list_prompts(self) -> list[MCPPrompt]:
|
||||
logger.debug("Handler called: list_prompts")
|
||||
|
||||
with fastmcp.server.context.Context(fastmcp=self):
|
||||
prompts = await self._middleware_list_prompts()
|
||||
return [prompt.to_mcp_prompt(name=prompt.key) for prompt in prompts]
|
||||
|
||||
async def _middleware_list_prompts(self) -> list[Prompt]:
|
||||
"""
|
||||
List all available prompts, in the format expected by the low-level MCP
|
||||
server.
|
||||
|
||||
"""
|
||||
logger.debug("Handler called: list_prompts")
|
||||
|
||||
async def _final_handler(
|
||||
async def _handler(
|
||||
context: MiddlewareContext[dict[str, Any]],
|
||||
) -> list[MCPPrompt]:
|
||||
prompts = await self.get_prompts()
|
||||
mcp_prompts: list[MCPPrompt] = []
|
||||
for key, prompt in prompts.items():
|
||||
) -> list[Prompt]:
|
||||
prompts = await self._list_prompts()
|
||||
|
||||
mcp_prompts: list[Prompt] = []
|
||||
for prompt in prompts:
|
||||
if self._should_enable_component(prompt):
|
||||
mcp_prompts.append(prompt.to_mcp_prompt(name=key))
|
||||
mcp_prompts.append(prompt)
|
||||
|
||||
return mcp_prompts
|
||||
|
||||
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
|
||||
|
|
@ -672,7 +705,42 @@ class FastMCP(Generic[LifespanResultT]):
|
|||
)
|
||||
|
||||
# Apply the middleware chain.
|
||||
return await self._apply_middleware(mw_context, _final_handler)
|
||||
return await self._apply_middleware(mw_context, _handler)
|
||||
|
||||
async def _list_prompts(self, apply_middleware: bool = True) -> list[Prompt]:
|
||||
"""
|
||||
List all available prompts.
|
||||
"""
|
||||
|
||||
if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
|
||||
prompts: list[Prompt] = []
|
||||
|
||||
# iterate such that new mounts overwrite older ones
|
||||
for mounted_server in self._mounted_servers:
|
||||
try:
|
||||
if apply_middleware:
|
||||
server_prompts = (
|
||||
await mounted_server.server._middleware_list_prompts()
|
||||
)
|
||||
else:
|
||||
server_prompts = await mounted_server.server._list_prompts()
|
||||
# Apply prefix to each prompt key if prefix exists
|
||||
if mounted_server.prefix:
|
||||
for prompt in server_prompts:
|
||||
prompt = prompt.with_key(
|
||||
f"{mounted_server.prefix}_{prompt.key}"
|
||||
)
|
||||
prompts.append(prompt)
|
||||
else:
|
||||
prompts.extend(server_prompts)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
|
||||
)
|
||||
continue
|
||||
prompts.extend(self._prompt_manager.get_prompts().values())
|
||||
self._cache.set("prompts", prompts)
|
||||
return prompts
|
||||
|
||||
async def _mcp_call_tool(
|
||||
self, key: str, arguments: dict[str, Any]
|
||||
|
|
|
|||
|
|
@ -21,9 +21,10 @@ class Recording:
|
|||
class RecordingMiddleware(MCPMiddleware):
|
||||
"""A middleware that automatically records all method calls."""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, name: str | None = None):
|
||||
super().__init__()
|
||||
self.calls: list[Recording] = []
|
||||
self.name = name
|
||||
|
||||
def __getattribute__(self, name: str) -> Callable:
|
||||
"""Dynamically create recording methods for any on_* method."""
|
||||
|
|
@ -95,7 +96,7 @@ class RecordingMiddleware(MCPMiddleware):
|
|||
@pytest.fixture
|
||||
def recording_middleware():
|
||||
"""Fixture that provides a recording middleware instance."""
|
||||
middleware = RecordingMiddleware()
|
||||
middleware = RecordingMiddleware(name="recording_middleware")
|
||||
yield middleware
|
||||
|
||||
|
||||
|
|
@ -215,7 +216,7 @@ class TestMiddlewareHooks:
|
|||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(
|
||||
method="resource-templates/list", times=3
|
||||
method="resources/templates/list", times=3
|
||||
)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
|
|
@ -234,3 +235,258 @@ class TestMiddlewareHooks:
|
|||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_list_prompts", times=1)
|
||||
|
||||
|
||||
class TestNestedMiddlewareHooks:
|
||||
@pytest.fixture
|
||||
@staticmethod
|
||||
def nested_middleware():
|
||||
return RecordingMiddleware(name="nested_middleware")
|
||||
|
||||
@pytest.fixture
|
||||
def nested_mcp_server(self, nested_middleware: RecordingMiddleware):
|
||||
mcp = FastMCP(name="Nested MCP")
|
||||
|
||||
@mcp.tool
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
@mcp.resource("resource://test")
|
||||
def test_resource() -> str:
|
||||
return "test resource"
|
||||
|
||||
@mcp.resource("resource://test-template/{x}")
|
||||
def test_resource_with_path(x: int) -> str:
|
||||
return f"test resource with {x}"
|
||||
|
||||
@mcp.prompt
|
||||
def test_prompt(x: str) -> str:
|
||||
return f"test prompt with {x}"
|
||||
|
||||
@mcp.tool
|
||||
async def progress_tool(context: Context) -> None:
|
||||
await context.report_progress(progress=1, total=10, message="test")
|
||||
|
||||
@mcp.tool
|
||||
async def log_tool(context: Context) -> None:
|
||||
await context.info(message="test log")
|
||||
|
||||
@mcp.tool
|
||||
async def sample_tool(context: Context) -> None:
|
||||
await context.sample("hello")
|
||||
|
||||
mcp.add_middleware(nested_middleware)
|
||||
|
||||
return mcp
|
||||
|
||||
async def test_call_tool_on_parent_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.call_tool("add", {"a": 1, "b": 2})
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="tools/call", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_call_tool", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=0)
|
||||
|
||||
async def test_call_tool_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.call_tool("nested_add", {"a": 1, "b": 2})
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="tools/call", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_call_tool", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(method="tools/call", times=3)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_call_tool", times=1)
|
||||
|
||||
async def test_read_resource_on_parent_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.read_resource("resource://test")
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="resources/read", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_read_resource", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=0)
|
||||
|
||||
async def test_read_resource_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.read_resource("resource://nested/test")
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="resources/read", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_read_resource", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(method="resources/read", times=3)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_read_resource", times=1)
|
||||
|
||||
async def test_get_prompt_on_parent_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.get_prompt("test_prompt", {"x": "test"})
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="prompts/get", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=0)
|
||||
|
||||
async def test_get_prompt_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.get_prompt("nested_test_prompt", {"x": "test"})
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="prompts/get", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(method="prompts/get", times=3)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_get_prompt", times=1)
|
||||
|
||||
async def test_list_tools_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.list_tools()
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="tools/list", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_list_tools", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(method="tools/list", times=3)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_list_tools", times=1)
|
||||
|
||||
async def test_list_resources_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.list_resources()
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(method="resources/list", times=3)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_list_resources", times=1)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(method="resources/list", times=3)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_list_resources", times=1)
|
||||
|
||||
async def test_list_resource_templates_on_nested_server(
|
||||
self,
|
||||
mcp_server: FastMCP,
|
||||
nested_mcp_server: FastMCP,
|
||||
recording_middleware: RecordingMiddleware,
|
||||
nested_middleware: RecordingMiddleware,
|
||||
):
|
||||
mcp_server.mount(nested_mcp_server, prefix="nested")
|
||||
|
||||
async with Client(mcp_server) as client:
|
||||
await client.list_resource_templates()
|
||||
|
||||
assert recording_middleware.assert_called(times=3)
|
||||
assert recording_middleware.assert_called(
|
||||
method="resources/templates/list", times=3
|
||||
)
|
||||
assert recording_middleware.assert_called(hook="on_message", times=1)
|
||||
assert recording_middleware.assert_called(hook="on_request", times=1)
|
||||
assert recording_middleware.assert_called(
|
||||
hook="on_list_resource_templates", times=1
|
||||
)
|
||||
|
||||
assert nested_middleware.assert_called(times=3)
|
||||
assert nested_middleware.assert_called(
|
||||
method="resources/templates/list", times=3
|
||||
)
|
||||
assert nested_middleware.assert_called(hook="on_message", times=1)
|
||||
assert nested_middleware.assert_called(hook="on_request", times=1)
|
||||
assert nested_middleware.assert_called(
|
||||
hook="on_list_resource_templates", times=1
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue