fastmcp/tests/server/test_provider_extensions.py
2026-10-02 20:51:00 -04:00

906 lines
33 KiB
Python

"""Provider-bundled extensions across composition, configuration, and startup."""
from __future__ import annotations
import logging
import threading
from collections.abc import AsyncIterator, Sequence
from contextlib import asynccontextmanager
from typing import Any, Literal
import mcp_types
import pytest
from mcp.server.context import ServerRequestContext
from mcp.shared.exceptions import MCPError
from mcp_types import METHOD_NOT_FOUND, RequestParams
from fastmcp import Client, FastMCP
from fastmcp.server.context import Context
from fastmcp.server.extensions import (
MethodBinding,
ServerExtension,
ToolCallContinuation,
ToolCallOutcome,
)
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
from fastmcp.server.providers import LocalProvider, Provider
from fastmcp.server.providers.aggregate import AggregateProvider
from fastmcp.server.providers.fastmcp_provider import FastMCPProvider
from fastmcp.server.transforms import Namespace
from fastmcp_tasks import TasksExtension
EXT_ID = "com.example/bundled"
class InspectRequest(mcp_types.Request):
method: Literal["bundled/inspect"] = "bundled/inspect"
params: RequestParams
class OptionalRequest(mcp_types.Request):
method: Literal["bundled/optional"] = "bundled/optional"
params: RequestParams
class InspectResult(mcp_types.Result):
server: str
tools: list[str]
class BundledExtension(ServerExtension):
identifier = EXT_ID
auto_register = True
def __init__(self, *, optional_method: bool = True) -> None:
self.optional_method = optional_method
self.lifecycle: list[str] = []
def settings(self) -> dict[str, Any]:
return {"optionalMethod": self.optional_method}
def methods(self) -> Sequence[MethodBinding]:
bindings = [MethodBinding("bundled/inspect", RequestParams, self.inspect)]
if self.optional_method:
bindings.append(
MethodBinding("bundled/optional", RequestParams, self.inspect)
)
return bindings
async def inspect(
self, ctx: ServerRequestContext[Any, Any], params: RequestParams
) -> InspectResult:
return InspectResult(
server=self.server.name,
tools=[tool.name for tool in await self.server.list_tools()],
)
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
self.lifecycle.append("enter")
try:
yield
finally:
self.lifecycle.append("exit")
class BundledProvider(Provider):
def __init__(self, extension: ServerExtension | None = None) -> None:
super().__init__()
self.extension = extension if extension is not None else BundledExtension()
def required_extensions(self) -> Sequence[ServerExtension]:
return [self.extension]
class AuditExtension(BundledExtension):
def __init__(self, *, enabled: bool = True) -> None:
super().__init__()
self.enabled = enabled
self.calls: list[str] = []
def settings(self) -> dict[str, Any]:
return {"enabled": self.enabled}
async def intercept_tool_call(
self,
params: mcp_types.CallToolRequestParams,
context: Context,
call_next: ToolCallContinuation,
) -> ToolCallOutcome:
if self.enabled:
self.calls.append(params.name)
return await call_next()
@pytest.mark.parametrize("construction", ["constructor", "add_provider"])
async def test_bundled_extension_serves_capability_and_method(construction: str):
provider = BundledProvider()
if construction == "constructor":
server = FastMCP("root", providers=[provider])
else:
server = FastMCP("root")
server.add_provider(provider)
async with Client(server, mode="auto") as client:
assert client.server_capabilities.extensions[EXT_ID] == {"optionalMethod": True}
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "root"
@pytest.mark.parametrize("wrapper", ["aggregate", "namespace", "nested"])
async def test_composed_providers_preserve_bundled_extensions(wrapper: str):
provider: Provider = BundledProvider()
if wrapper == "aggregate":
provider = AggregateProvider([provider])
elif wrapper == "namespace":
provider = provider.wrap_transform(Namespace("ns"))
else:
provider = AggregateProvider(
[AggregateProvider([provider.wrap_transform(Namespace("inner"))])]
).wrap_transform(Namespace("outer"))
server = FastMCP("root", providers=[provider])
async with Client(server, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
@pytest.mark.parametrize("attachment", ["mount", "add_provider", "adapter"])
async def test_mounted_extensions_use_root_registry_and_preserve_child(attachment: str):
child = FastMCP("child")
extension = BundledExtension(optional_method=False)
child.add_extension(extension)
@child.tool
def greet() -> str:
return "hello"
middle = FastMCP("middle")
if attachment == "mount":
middle.mount(child, namespace="child", tool_names={"greet": "hello"})
else:
provider = child if attachment == "add_provider" else FastMCPProvider(child)
middle.add_provider(provider, namespace="child")
root = FastMCP("root")
root.mount(middle, namespace="middle")
assert extension.server is child
assert root._extensions[EXT_ID] is not middle._extensions[EXT_ID]
async with Client(root, mode="auto") as client:
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "root"
expected = "hello" if attachment == "mount" else "greet"
assert result.tools == [f"middle_child_{expected}"]
assert client.server_capabilities.extensions[EXT_ID] == {
"optionalMethod": False
}
async with Client(child, mode="auto") as client:
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "child"
assert result.tools == ["greet"]
def test_one_provider_can_be_used_by_independent_servers():
provider = BundledProvider()
first = FastMCP("first", providers=[provider])
second = FastMCP("second", providers=[provider])
assert first._extensions[EXT_ID].server is first
assert second._extensions[EXT_ID].server is second
assert first._extensions[EXT_ID] is not second._extensions[EXT_ID]
with pytest.raises(RuntimeError, match="not bound"):
_ = provider.extension.server
@pytest.mark.parametrize("explicit_first", [True, False])
async def test_explicit_configuration_wins_and_removes_optional_method(
explicit_first: bool,
):
server = FastMCP("root")
explicit = BundledExtension(optional_method=False)
if explicit_first:
server.add_extension(explicit)
server.add_provider(BundledProvider())
if not explicit_first:
server.add_extension(explicit)
async with Client(server, mode="auto") as client:
assert client.server_capabilities.extensions[EXT_ID] == {
"optionalMethod": False
}
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "root"
with pytest.raises(MCPError) as exc:
await client.session.send_request(
OptionalRequest(params=RequestParams()), InspectResult
)
assert exc.value.error.code == METHOD_NOT_FOUND
assert server._extensions[EXT_ID] is explicit
assert explicit.lifecycle == ["enter", "exit"]
def test_duplicate_explicit_registration_still_raises():
server = FastMCP("root", providers=[BundledProvider()])
server.add_extension(BundledExtension())
with pytest.raises(ValueError, match="already registered"):
server.add_extension(BundledExtension())
def test_same_bundled_configuration_is_deduplicated(caplog: pytest.LogCaptureFixture):
server = FastMCP("root")
with caplog.at_level(logging.WARNING):
server.add_provider(AggregateProvider([BundledProvider(), BundledProvider()]))
assert len(server._extensions) == 1
assert "Conflicting" not in caplog.text
async def test_conflicting_settings_warn_once_and_keep_first(
caplog: pytest.LogCaptureFixture,
):
server = FastMCP("root")
with caplog.at_level(logging.WARNING):
server.add_provider(BundledProvider(BundledExtension(optional_method=False)))
first = server._extensions[EXT_ID]
server.add_provider(BundledProvider())
async with Client(server, mode="auto"):
assert server._extensions[EXT_ID] is first
messages = [r.message for r in caplog.records if "Conflicting" in r.message]
assert len(messages) == 1
assert "add_extension()" in messages[0]
def test_explicit_configuration_suppresses_conflict_warning(
caplog: pytest.LogCaptureFixture,
):
server = FastMCP("root")
server.add_extension(BundledExtension(optional_method=False))
with caplog.at_level(logging.WARNING):
server.add_provider(BundledProvider())
assert "Conflicting" not in caplog.text
def test_registration_debug_log_names_extension_provider_and_server(
caplog: pytest.LogCaptureFixture,
):
with caplog.at_level(logging.DEBUG, logger="fastmcp.server.mixins.extensions"):
FastMCP("receiving-server", providers=[BundledProvider()])
assert EXT_ID in caplog.text
assert "BundledProvider" in caplog.text
assert "receiving-server" in caplog.text
assert "automatically" in caplog.text
def test_explicit_precedence_is_logged(caplog: pytest.LogCaptureFixture):
server = FastMCP("root")
server.add_extension(BundledExtension(optional_method=False))
with caplog.at_level(logging.DEBUG, logger="fastmcp.server.mixins.extensions"):
server.add_provider(BundledProvider())
assert "Using explicitly registered extension" in caplog.text
async def test_extensions_added_to_mounted_child_before_startup_are_discovered():
child = FastMCP("child")
root = FastMCP("root")
root.mount(child)
child.add_provider(BundledProvider())
async with Client(root, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
root_extension = root._extensions[EXT_ID]
assert isinstance(root_extension, BundledExtension)
assert root_extension.lifecycle == ["enter"]
child_extension = child._extensions[EXT_ID]
assert isinstance(child_extension, BundledExtension)
assert child_extension.lifecycle == []
assert root_extension.lifecycle == ["enter", "exit"]
async def test_extensions_added_to_aggregate_before_startup_are_discovered():
aggregate = AggregateProvider()
root = FastMCP("root", providers=[aggregate])
aggregate.add_provider(BundledProvider())
async with Client(root, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
async def test_provider_added_in_user_lifespan_is_discovered():
@asynccontextmanager
async def lifespan(server: FastMCP) -> AsyncIterator[None]:
server.add_provider(BundledProvider())
yield
root = FastMCP("root", lifespan=lifespan)
async with Client(root, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
async def test_new_extension_after_startup_rejected_without_adding_provider():
root = FastMCP("root")
async with Client(root, mode="auto"):
with pytest.raises(RuntimeError, match="lifespan has already started"):
root.add_provider(BundledProvider())
assert len(root.providers) == 1
assert EXT_ID not in root._extensions
async def test_explicit_override_after_startup_rejected():
root = FastMCP("root", providers=[BundledProvider()])
original = root._extensions[EXT_ID]
async with Client(root, mode="auto"):
with pytest.raises(RuntimeError, match="lifespan has already started"):
root.add_extension(BundledExtension(optional_method=False))
assert root._extensions[EXT_ID] is original
async def test_registration_during_extension_startup_rejected():
class RegisteringExtension(ServerExtension):
identifier = "com.example/registering"
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
with pytest.raises(RuntimeError, match="lifespan has already started"):
self.server.add_extension(BundledExtension())
yield
root = FastMCP("root")
root.add_extension(RegisteringExtension())
async with Client(root, mode="auto"):
assert EXT_ID not in root._extensions
@pytest.mark.parametrize("bundled", [False, True])
def test_extension_auto_registration_requires_opt_in(bundled: bool):
class ExplicitExtension(ServerExtension):
identifier = EXT_ID
root = FastMCP("root")
if bundled:
with pytest.raises(ValueError, match="does not allow automatic registration"):
root.add_provider(BundledProvider(ExplicitExtension()))
assert len(root.providers) == 1
else:
child = FastMCP("child")
child.add_extension(ExplicitExtension())
root.mount(child)
assert EXT_ID not in root._extensions
def test_non_automatic_dependency_can_be_satisfied_explicitly():
class ExplicitExtension(ServerExtension):
identifier = EXT_ID
root = FastMCP("root")
root.add_extension(ExplicitExtension())
root.add_provider(BundledProvider(ExplicitExtension()))
def test_tasks_extension_does_not_propagate():
child = FastMCP("child")
child.add_extension(TasksExtension())
root = FastMCP("root")
root.mount(child)
assert TasksExtension.identifier not in root._extensions
def test_cloning_clears_binding_and_isolates_mutable_state():
child = FastMCP("child")
extension = BundledExtension(optional_method=False)
child.add_extension(extension)
clone = extension.clone()
assert clone.settings() == extension.settings()
assert isinstance(clone, BundledExtension)
clone.lifecycle.append("changed")
assert extension.lifecycle == []
assert extension.server is child
with pytest.raises(RuntimeError, match="not bound"):
_ = clone.server
def test_registration_rejects_sharing_a_bound_instance():
extension = BundledExtension()
child = FastMCP("child")
child.add_extension(extension)
root = FastMCP("root")
with pytest.raises(ValueError, match="already bound to another server"):
root.add_extension(extension)
assert extension.server is child
def test_unrelated_extension_method_collision_is_rejected():
class OtherExtension(BundledExtension):
identifier = "com.example/other"
root = FastMCP("root")
root.add_extension(OtherExtension())
with pytest.raises(
ValueError, match="method 'bundled/inspect' is already registered"
):
root.add_provider(BundledProvider())
assert EXT_ID not in root._extensions
assert len(root.providers) == 1
@pytest.mark.parametrize("failure", ["collision", "methods"])
async def test_rejected_explicit_extension_can_be_registered_on_another_server(
failure: str,
):
class RecoverableExtension(BundledExtension):
identifier = "com.example/recoverable"
def methods(self) -> Sequence[MethodBinding]:
if failure == "methods" and self.server.name == "first":
raise RuntimeError("Cannot build methods")
return super().methods()
first = FastMCP("first")
if failure == "collision":
first.add_extension(BundledExtension())
extension = RecoverableExtension()
error = ValueError if failure == "collision" else RuntimeError
message = "already registered" if failure == "collision" else "Cannot build methods"
with pytest.raises(error, match=message):
first.add_extension(extension)
assert extension.identifier not in first._extensions
with pytest.raises(RuntimeError, match="not bound"):
_ = extension.server
second = FastMCP("second")
second.add_extension(extension)
async with Client(second, mode="auto") as client:
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "second"
assert extension.server is second
def test_failed_registration_preserves_an_existing_server_binding():
class FailingExtension(BundledExtension):
fail_methods = False
def methods(self) -> Sequence[MethodBinding]:
if self.fail_methods:
raise RuntimeError("Cannot build methods")
return super().methods()
root = FastMCP("root", providers=[BundledProvider(FailingExtension())])
extension = root._extensions[EXT_ID]
assert isinstance(extension, FailingExtension)
handler = root._mcp_server.get_request_handler("bundled/inspect")
extension.fail_methods = True
with pytest.raises(RuntimeError, match="Cannot build methods"):
root.add_extension(extension)
assert extension.server is root
assert root._extensions[EXT_ID] is extension
assert root._mcp_server.get_request_handler("bundled/inspect") is handler
async def test_failed_bundle_rolls_back_all_extensions_and_handlers():
class ExplicitExtension(ServerExtension):
identifier = "com.example/explicit"
root = FastMCP("root")
bundle = AggregateProvider(
[BundledProvider(), BundledProvider(ExplicitExtension())]
)
with pytest.raises(ValueError, match="does not allow automatic registration"):
root.add_provider(bundle)
assert root._extensions == {}
assert root._auto_extensions == set()
assert root._mcp_server.get_request_handler("bundled/inspect") is None
assert len(root.providers) == 1
# A corrected attempt must work without stale methods or duplicate IDs.
root.add_extension(ExplicitExtension())
root.add_provider(bundle)
async with Client(root, mode="auto") as client:
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "root"
def test_settings_can_depend_on_receiving_server(caplog: pytest.LogCaptureFixture):
class ServerAwareExtension(BundledExtension):
def settings(self) -> dict[str, Any]:
return {"server": self.server.name}
child = FastMCP("child", providers=[BundledProvider(ServerAwareExtension())])
root = FastMCP("root")
with caplog.at_level(logging.WARNING):
root.mount(child)
root.add_provider(BundledProvider(ServerAwareExtension()))
assert root._extensions[EXT_ID].settings() == {"server": "root"}
assert "Conflicting" not in caplog.text
def test_custom_clone_can_reconstruct_non_copyable_configuration():
class LockedExtension(BundledExtension):
def __init__(self, *, optional_method: bool = True) -> None:
super().__init__(optional_method=optional_method)
self.lock = threading.Lock()
def clone(self) -> ServerExtension:
return LockedExtension(optional_method=self.optional_method)
extension = LockedExtension(optional_method=False)
root = FastMCP("root", providers=[BundledProvider(extension)])
registered = root._extensions[EXT_ID]
assert isinstance(registered, LockedExtension)
assert registered.lock is not extension.lock
assert registered.settings() == {"optionalMethod": False}
@pytest.mark.parametrize("invalid_clone", ["same_instance", "wrong_identifier"])
def test_invalid_clone_rejected_without_rebinding_source(invalid_clone: str):
class InvalidCloneExtension(BundledExtension):
def clone(self) -> ServerExtension:
if invalid_clone == "same_instance":
return self
clone = BundledExtension()
clone.identifier = "com.example/wrong"
return clone
child = FastMCP("child")
extension = InvalidCloneExtension()
child.add_extension(extension)
root = FastMCP("root")
with pytest.raises(ValueError, match="separate extension with the same identifier"):
root.mount(child)
assert extension.server is child
assert root._extensions == {}
async def test_explicit_child_extension_added_after_mount_is_discovered():
child = FastMCP("child")
root = FastMCP("root")
root.mount(child)
child.add_extension(BundledExtension(optional_method=False))
async with Client(root, mode="auto") as client:
assert client.server_capabilities.extensions[EXT_ID] == {
"optionalMethod": False
}
async def test_explicit_root_configuration_wins_over_mounted_child():
child = FastMCP("child", providers=[BundledProvider()])
root = FastMCP("root")
root.mount(child)
root.add_extension(BundledExtension(optional_method=False))
async with Client(root, mode="auto") as client:
assert client.server_capabilities.extensions[EXT_ID] == {
"optionalMethod": False
}
assert child._extensions[EXT_ID].settings() == {"optionalMethod": True}
async def test_mounted_child_lifespan_cannot_introduce_a_new_root_extension():
@asynccontextmanager
async def lifespan(server: FastMCP) -> AsyncIterator[None]:
server.add_provider(BundledProvider())
yield
child = FastMCP("child", lifespan=lifespan)
root = FastMCP("root")
root.mount(child)
with pytest.raises(RuntimeError, match="unavailable in a running server"):
async with root._lifespan_manager():
pass
assert child._extensions == {}
assert len(child.providers) == 1
async def test_mounted_child_lifespan_can_use_an_existing_root_extension():
@asynccontextmanager
async def lifespan(server: FastMCP) -> AsyncIterator[None]:
server.add_provider(BundledProvider())
yield
child = FastMCP("child", lifespan=lifespan)
root = FastMCP("root", providers=[BundledProvider()])
root.mount(child)
async with Client(root, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
@pytest.mark.parametrize("wrapped", [False, True])
async def test_running_aggregate_rejects_new_extension_without_exposing_components(
wrapped: bool,
):
class BundledLocalProvider(LocalProvider):
def required_extensions(self) -> Sequence[ServerExtension]:
return [BundledExtension()]
provider = BundledLocalProvider()
@provider.tool
def marker() -> str:
return "marker"
aggregate = AggregateProvider()
composed: Provider = aggregate
if wrapped:
composed = AggregateProvider([aggregate.wrap_transform(Namespace("ns"))])
root = FastMCP("root", providers=[composed])
async with Client(root, mode="auto") as client:
with pytest.raises(RuntimeError, match="unavailable in a running server"):
aggregate.add_provider(provider)
assert await client.list_tools() == []
assert root._extensions == {}
assert aggregate.providers == []
async def test_running_aggregate_accepts_bundle_supported_by_root():
aggregate = AggregateProvider()
root = FastMCP("root", providers=[aggregate])
root.add_extension(BundledExtension())
async with Client(root, mode="auto"):
aggregate.add_provider(BundledProvider())
assert len(aggregate.providers) == 1
async def test_running_aggregate_accepts_components_without_new_extensions():
aggregate = AggregateProvider()
root = FastMCP("root", providers=[aggregate])
provider = LocalProvider()
@provider.tool
def marker() -> str:
return "marker"
async with Client(root, mode="auto") as client:
aggregate.add_provider(provider)
assert [tool.name for tool in await client.list_tools()] == ["marker"]
async def test_shared_aggregate_checks_every_active_server_and_releases_scopes():
aggregate = AggregateProvider()
first = FastMCP("first", providers=[aggregate])
first.add_extension(BundledExtension())
second = FastMCP("second", providers=[aggregate])
async with Client(first, mode="auto"):
async with Client(second, mode="auto"):
with pytest.raises(RuntimeError, match="unavailable in a running server"):
aggregate.add_provider(BundledProvider())
aggregate.add_provider(BundledProvider())
assert aggregate._extension_scopes == []
@pytest.mark.parametrize("enabled", [True, False])
async def test_propagated_interceptor_runs_once_with_root_configuration(enabled: bool):
child = FastMCP("child", providers=[BundledProvider(AuditExtension())])
@child.tool
def marker() -> str:
return "marker"
middle = FastMCP("middle")
middle.mount(child, namespace="child", tool_names={"marker": "renamed"})
root = FastMCP("root")
root.mount(middle, namespace="middle")
explicit = AuditExtension(enabled=enabled)
root.add_extension(explicit)
child_extension = child._extensions[EXT_ID]
middle_extension = middle._extensions[EXT_ID]
assert isinstance(child_extension, AuditExtension)
assert isinstance(middle_extension, AuditExtension)
async with Client(root, mode="auto") as client:
await client.call_tool("middle_child_renamed")
assert explicit.calls == (["middle_child_renamed"] if enabled else [])
assert child_extension.calls == []
assert middle_extension.calls == []
async with Client(child, mode="auto") as client:
await client.call_tool("marker")
assert child_extension.calls == ["marker"]
async def test_delegation_preserves_programmatic_child_calls():
child = FastMCP("child", providers=[BundledProvider(AuditExtension())])
@child.tool
def inner() -> str:
return "inner"
@child.tool
async def outer() -> str:
await child.call_tool("inner")
return "outer"
root = FastMCP("root")
root.mount(child)
async with Client(root, mode="auto") as client:
await client.call_tool("outer")
root_extension = root._extensions[EXT_ID]
child_extension = child._extensions[EXT_ID]
assert isinstance(root_extension, AuditExtension)
assert isinstance(child_extension, AuditExtension)
assert root_extension.calls == ["outer"]
assert child_extension.calls == ["inner"]
async def test_non_propagating_interceptors_keep_server_local_behavior():
class LocalAuditExtension(AuditExtension):
auto_register = False
child = FastMCP("child")
child_extension = LocalAuditExtension()
child.add_extension(child_extension)
@child.tool
def marker() -> str:
return "marker"
root = FastMCP("root")
root_extension = LocalAuditExtension()
root.add_extension(root_extension)
root.mount(child)
async with Client(root, mode="auto") as client:
await client.call_tool("marker")
assert root_extension.calls == ["marker"]
assert child_extension.calls == ["marker"]
async def test_nested_child_setup_preserves_ancestor_interceptor_configuration():
@asynccontextmanager
async def lifespan(server: FastMCP) -> AsyncIterator[None]:
server.add_provider(BundledProvider(AuditExtension()))
yield
child = FastMCP("child", lifespan=lifespan)
@child.tool
def marker() -> str:
return "marker"
middle = FastMCP("middle")
middle.mount(child)
root = FastMCP("root")
root.mount(middle)
root.add_extension(AuditExtension(enabled=False))
async with Client(root, mode="auto") as client:
await client.call_tool("marker")
child_extension = child._extensions[EXT_ID]
assert isinstance(child_extension, AuditExtension)
assert child_extension.calls == []
assert middle._extensions == {}
async def test_delegation_preserves_programmatic_calls_from_child_middleware():
child = FastMCP("child", providers=[BundledProvider(AuditExtension())])
@child.tool
def inner() -> str:
return "inner"
@child.tool
def outer() -> str:
return "outer"
class CallingMiddleware(Middleware):
async def on_call_tool(
self,
context: MiddlewareContext[mcp_types.CallToolRequestParams],
call_next: CallNext[mcp_types.CallToolRequestParams, Any],
) -> Any:
if context.message.name == "outer":
await child.call_tool("inner")
return await call_next(context)
child.add_middleware(CallingMiddleware())
root = FastMCP("root")
root.mount(child)
async with Client(root, mode="auto") as client:
await client.call_tool("outer")
root_extension = root._extensions[EXT_ID]
child_extension = child._extensions[EXT_ID]
assert isinstance(root_extension, AuditExtension)
assert isinstance(child_extension, AuditExtension)
assert root_extension.calls == ["outer"]
assert child_extension.calls == ["inner"]
async def test_aggregate_setup_cannot_introduce_a_new_root_extension():
aggregate = AggregateProvider()
class SetupProvider(Provider):
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
aggregate.add_provider(BundledProvider())
yield
aggregate.add_provider(SetupProvider())
root = FastMCP("root", providers=[aggregate])
with pytest.raises(RuntimeError, match="unavailable in a running server"):
async with root._lifespan_manager():
pass
assert root._extensions == {}
assert len(aggregate.providers) == 1
class MutatingProvider(Provider):
"""A provider whose lifespan adds a bundled provider to a later sibling."""
def __init__(self, target: FastMCP | AggregateProvider) -> None:
super().__init__()
self.target = target
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
self.target.add_provider(BundledProvider())
yield
@pytest.mark.parametrize("target_kind", ["mounted_server", "aggregate"])
async def test_lifespan_cannot_introduce_a_root_extension_via_a_later_sibling(
target_kind: Literal["mounted_server", "aggregate"],
):
target = (
FastMCP("later") if target_kind == "mounted_server" else AggregateProvider()
)
root = FastMCP("root", providers=[MutatingProvider(target)])
if isinstance(target, FastMCP):
root.mount(target)
else:
root.add_provider(target)
with pytest.raises(RuntimeError, match="lifespan has already started"):
async with root._lifespan_manager():
pass
assert root._extensions == {}
async def test_lifespan_can_add_a_bundle_the_root_already_supports_to_a_later_sibling():
later = FastMCP("later")
root = FastMCP("root", providers=[BundledProvider(), MutatingProvider(later)])
root.mount(later)
async with Client(root, mode="auto") as client:
assert EXT_ID in client.server_capabilities.extensions
@pytest.mark.parametrize("target_kind", ["mounted_server", "aggregate"])
@pytest.mark.parametrize("wrapped", [False, True])
async def test_aggregate_startup_rejects_new_extension_on_a_later_descendant(
target_kind: Literal["mounted_server", "aggregate"], wrapped: bool
):
later = FastMCP("later") if target_kind == "mounted_server" else AggregateProvider()
aggregate = AggregateProvider([MutatingProvider(later)])
aggregate.add_provider(later)
composed: Provider = aggregate
if wrapped:
composed = AggregateProvider([aggregate.wrap_transform(Namespace("ns"))])
root = FastMCP("root", providers=[composed])
with pytest.raises(RuntimeError, match="unavailable in a running server"):
async with root._lifespan_manager():
pass
assert root._extensions == {}
assert aggregate._extension_scopes == []
@pytest.mark.parametrize("target_kind", ["mounted_server", "aggregate"])
@pytest.mark.parametrize("wrapped", [False, True])
async def test_aggregate_startup_accepts_supported_bundle_on_a_later_descendant(
target_kind: Literal["mounted_server", "aggregate"], wrapped: bool
):
later = FastMCP("later") if target_kind == "mounted_server" else AggregateProvider()
aggregate = AggregateProvider([MutatingProvider(later)])
aggregate.add_provider(later)
composed: Provider = aggregate
if wrapped:
composed = AggregateProvider([aggregate.wrap_transform(Namespace("ns"))])
root = FastMCP("root", providers=[BundledProvider(), composed])
async with Client(root, mode="auto") as client:
result = await client.session.send_request(
InspectRequest(params=RequestParams()), InspectResult
)
assert result.server == "root"
extension = root._extensions[EXT_ID]
assert isinstance(extension, BundledExtension)
assert extension.lifecycle == ["enter"]
assert extension.lifecycle == ["enter", "exit"]
assert aggregate._extension_scopes == []