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

278 lines
9.9 KiB
Python

"""Extension requirements follow every runtime using live provider composition."""
from collections.abc import AsyncIterator, Sequence
from contextlib import AsyncExitStack, asynccontextmanager
import anyio
import pytest
from fastmcp import Client, FastMCP
from fastmcp.server.extensions import ServerExtension
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
EXT_ID = "com.example/runtime"
class RuntimeExtension(ServerExtension):
identifier = EXT_ID
auto_register = True
def __init__(self) -> None:
self.events: list[str] = []
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
self.events.append("enter")
try:
yield
finally:
self.events.append("exit")
class BundledLocalProvider(LocalProvider):
def required_extensions(self) -> Sequence[ServerExtension]:
return [RuntimeExtension()]
@pytest.fixture
def bundle() -> BundledLocalProvider:
provider = BundledLocalProvider()
@provider.tool
def marker() -> str:
return "marker"
return provider
@pytest.mark.parametrize("wrapped", [False, True])
async def test_running_mounted_child_accepts_root_supported_bundle(
bundle: BundledLocalProvider, wrapped: bool
):
child = FastMCP("child")
root = FastMCP("root")
extension = RuntimeExtension()
root.add_extension(extension)
provider: Provider = FastMCPProvider(child)
if wrapped:
provider = provider.wrap_transform(Namespace("ns"))
root.add_provider(provider)
async with Client(root, mode="auto") as client:
with pytest.raises(RuntimeError, match="lifespan has already started"):
child.add_extension(RuntimeExtension())
child.add_provider(bundle)
name = "ns_marker" if wrapped else "marker"
assert [tool.name for tool in await client.list_tools()] == [name]
assert (await client.call_tool(name)).data == "marker"
assert EXT_ID in client.server_capabilities.extensions
assert extension.events == ["enter"]
registered = child._extensions[EXT_ID]
assert isinstance(registered, RuntimeExtension)
assert registered.events == []
assert extension.events == ["enter", "exit"]
assert child._extension_scopes == []
async def test_child_provider_lifespan_accepts_root_supported_bundle(
bundle: BundledLocalProvider,
):
child = FastMCP("child")
class SetupProvider(Provider):
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
child.add_provider(bundle)
yield
child.add_provider(SetupProvider())
root = FastMCP("root")
root.add_extension(RuntimeExtension())
root.mount(child)
async with Client(root, mode="auto") as client:
assert (await client.call_tool("marker")).data == "marker"
@pytest.mark.parametrize("wrapped", [False, True])
async def test_shared_mounted_aggregate_checks_every_root(
bundle: BundledLocalProvider, wrapped: bool
):
aggregate = AggregateProvider()
provider: Provider = aggregate
if wrapped:
provider = AggregateProvider([aggregate.wrap_transform(Namespace("ns"))])
child = FastMCP("shared child", providers=[provider])
first = FastMCP("first")
first.add_extension(RuntimeExtension())
first.mount(child)
second = FastMCP("second")
second.mount(child)
async with Client(first, mode="auto"):
async with Client(second, mode="auto") as client:
with pytest.raises(RuntimeError, match="unavailable in a running server"):
aggregate.add_provider(bundle)
assert aggregate.providers == []
assert await client.list_tools() == []
assert EXT_ID not in client.server_capabilities.extensions
# Releasing the unsupported root must release its restriction as well.
aggregate.add_provider(bundle)
assert aggregate._extension_scopes == []
assert child._extension_scopes == []
async def test_shared_mounted_child_accepts_bundle_supported_by_all_roots(
bundle: BundledLocalProvider,
):
child = FastMCP("shared child")
first, second = FastMCP("first"), FastMCP("second")
for root in (first, second):
root.add_extension(RuntimeExtension())
root.mount(child)
async with Client(first, mode="auto") as first_client:
async with Client(second, mode="auto") as second_client:
child.add_provider(bundle)
assert (await first_client.call_tool("marker")).data == "marker"
assert (await second_client.call_tool("marker")).data == "marker"
assert child._extension_scopes == []
async def test_shared_child_keeps_second_roots_restrictions_after_first_exits(
bundle: BundledLocalProvider,
):
aggregate = AggregateProvider()
child = FastMCP("shared child", providers=[aggregate])
first, second = FastMCP("first"), FastMCP("second")
first.add_extension(RuntimeExtension())
first.mount(child)
second.mount(child)
first_running, stop_first, first_stopped = (
anyio.Event(),
anyio.Event(),
anyio.Event(),
)
async def serve_first() -> None:
async with Client(first, mode="auto"):
first_running.set()
await stop_first.wait()
first_stopped.set()
async with anyio.create_task_group() as tasks:
tasks.start_soon(serve_first)
await first_running.wait()
async with Client(second, mode="auto") as client:
stop_first.set()
await first_stopped.wait()
with pytest.raises(RuntimeError, match="unavailable in a running server"):
aggregate.add_provider(bundle)
assert await client.list_tools() == []
assert aggregate._extension_scopes == []
async def test_shared_child_rejection_leaves_its_registrations_unchanged(
bundle: BundledLocalProvider,
):
child = FastMCP("shared child")
first, second = FastMCP("first"), FastMCP("second")
first.add_extension(RuntimeExtension())
first.mount(child)
second.mount(child)
async with Client(first, mode="auto"):
async with Client(second, mode="auto"):
with pytest.raises(RuntimeError):
child.add_provider(bundle)
assert child._extensions == {}
assert len(child.providers) == 1
child.add_provider(bundle)
async def test_dynamically_attached_descendant_inherits_all_root_restrictions(
bundle: BundledLocalProvider,
):
aggregate = AggregateProvider()
child = FastMCP("shared child", providers=[aggregate])
first, second = FastMCP("first"), FastMCP("second")
first.add_extension(RuntimeExtension())
first.mount(child)
second.mount(child)
descendant = AggregateProvider()
async with Client(first, mode="auto"):
async with Client(second, mode="auto") as client:
aggregate.add_provider(descendant.wrap_transform(Namespace("ns")))
with pytest.raises(RuntimeError, match="unavailable in a running server"):
descendant.add_provider(bundle)
assert await client.list_tools() == []
descendant.add_provider(bundle)
assert descendant._extension_scopes == []
@pytest.mark.parametrize("standalone_first", [False, True])
async def test_standalone_child_still_rejects_unavailable_extension(
bundle: BundledLocalProvider, standalone_first: bool
):
child = FastMCP("child")
root = FastMCP("root")
root.add_extension(RuntimeExtension())
root.mount(child)
first, second = (child, root) if standalone_first else (root, child)
async with Client(first, mode="auto") as first_client:
async with Client(second, mode="auto") as second_client:
with pytest.raises(RuntimeError):
child.add_provider(bundle)
standalone = first_client if standalone_first else second_client
assert await standalone.list_tools() == []
assert child._extensions == {}
async def test_standalone_aggregate_releases_descendant_restrictions(
bundle: BundledLocalProvider,
):
descendant = AggregateProvider()
aggregate = AggregateProvider([descendant])
async with aggregate.lifespan():
with pytest.raises(RuntimeError, match="unavailable in a running server"):
descendant.add_provider(bundle)
descendant.add_provider(bundle)
assert aggregate._extension_scopes == []
assert descendant._extension_scopes == []
async def test_overlapping_standalone_aggregate_lifespans_keep_independent_scopes(
bundle: BundledLocalProvider,
):
descendant = AggregateProvider()
aggregate = AggregateProvider([descendant])
async with AsyncExitStack() as first, AsyncExitStack() as second:
await first.enter_async_context(aggregate.lifespan())
await second.enter_async_context(aggregate.lifespan())
await first.aclose()
with pytest.raises(RuntimeError, match="unavailable in a running server"):
descendant.add_provider(bundle)
descendant.add_provider(bundle)
assert aggregate._extension_scopes == []
assert descendant._extension_scopes == []
async def test_failed_startup_releases_extension_restrictions(
bundle: BundledLocalProvider,
):
aggregate = AggregateProvider()
child = FastMCP("child", providers=[aggregate])
class FailingProvider(Provider):
@asynccontextmanager
async def lifespan(self) -> AsyncIterator[None]:
raise RuntimeError("setup failed")
yield # pragma: no cover
root = FastMCP("root", providers=[child, FailingProvider()])
with pytest.raises(RuntimeError, match="setup failed"):
async with root._lifespan_manager():
pass
assert child._extension_scopes == []
assert aggregate._extension_scopes == []
aggregate.add_provider(bundle)