mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-10-09 22:43:20 +02:00
906 lines
33 KiB
Python
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 == []
|