mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-24 06:24:18 +02:00
Merge pull request #2901 from jlowin/refactor-provider-inheritance-v2
This commit is contained in:
commit
5d12afdbf6
15 changed files with 601 additions and 634 deletions
|
|
@ -188,10 +188,14 @@ class ComponentService:
|
|||
if resource_keys:
|
||||
self._server.enable(keys=resource_keys)
|
||||
resource = await self._server.get_resource(uri)
|
||||
if resource is None:
|
||||
raise NotFoundError(f"Resource {uri!r} not found after enabling")
|
||||
return resource
|
||||
if template_keys:
|
||||
self._server.enable(keys=template_keys)
|
||||
template = await self._server.get_resource_template(uri)
|
||||
if template is None:
|
||||
raise NotFoundError(f"Template {uri!r} not found after enabling")
|
||||
return template
|
||||
|
||||
# 2. Check mounted servers via FastMCPProvider
|
||||
|
|
@ -234,11 +238,15 @@ class ComponentService:
|
|||
if resource_keys:
|
||||
# Get the highest version to return before disabling
|
||||
resource = await self._server.get_resource(uri)
|
||||
if resource is None:
|
||||
raise NotFoundError(f"Resource {uri!r} not found")
|
||||
self._server.disable(keys=resource_keys)
|
||||
return resource
|
||||
if template_keys:
|
||||
# Get the highest version to return before disabling
|
||||
template = await self._server.get_resource_template(uri)
|
||||
if template is None:
|
||||
raise NotFoundError(f"Template {uri!r} not found")
|
||||
self._server.disable(keys=template_keys)
|
||||
return template
|
||||
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ from collections.abc import Sequence
|
|||
|
||||
import mcp.types as mt
|
||||
|
||||
from fastmcp.exceptions import AuthorizationError, NotFoundError
|
||||
from fastmcp.exceptions import AuthorizationError
|
||||
from fastmcp.prompts.prompt import Prompt, PromptResult
|
||||
from fastmcp.resources.resource import Resource, ResourceResult
|
||||
from fastmcp.resources.template import ResourceTemplate
|
||||
|
|
@ -136,11 +136,16 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for tool '{tool_name}': missing context"
|
||||
)
|
||||
|
||||
tool = await fastmcp.fastmcp.get_tool(tool_name)
|
||||
# Get tool (component auth is checked in get_tool, raises if unauthorized)
|
||||
tool = await fastmcp.fastmcp._get_tool(tool_name)
|
||||
if tool is None:
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for tool '{tool_name}': tool not found"
|
||||
)
|
||||
|
||||
# Global auth check
|
||||
token = get_access_token()
|
||||
ctx = AuthContext(token=token, component=tool)
|
||||
|
||||
if not run_auth_checks(self.auth, ctx):
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for tool '{tool_name}': insufficient permissions"
|
||||
|
|
@ -196,15 +201,18 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for resource '{uri}': missing context"
|
||||
)
|
||||
|
||||
# Try concrete resource first, then template (for template-backed URIs)
|
||||
try:
|
||||
component = await fastmcp.fastmcp.get_resource(str(uri))
|
||||
except NotFoundError:
|
||||
component = await fastmcp.fastmcp.get_resource_template(str(uri))
|
||||
# Get resource/template (component auth is checked in get_*, raises if unauthorized)
|
||||
component = await fastmcp.fastmcp._get_resource(str(uri))
|
||||
if component is None:
|
||||
component = await fastmcp.fastmcp._get_resource_template(str(uri))
|
||||
if component is None:
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for resource '{uri}': resource not found"
|
||||
)
|
||||
|
||||
# Global auth check
|
||||
token = get_access_token()
|
||||
ctx = AuthContext(token=token, component=component)
|
||||
|
||||
if not run_auth_checks(self.auth, ctx):
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for resource '{uri}': insufficient permissions"
|
||||
|
|
@ -286,11 +294,16 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for prompt '{prompt_name}': missing context"
|
||||
)
|
||||
|
||||
prompt = await fastmcp.fastmcp.get_prompt(prompt_name)
|
||||
# Get prompt (component auth is checked in get_prompt, raises if unauthorized)
|
||||
prompt = await fastmcp.fastmcp._get_prompt(prompt_name)
|
||||
if prompt is None:
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for prompt '{prompt_name}': prompt not found"
|
||||
)
|
||||
|
||||
# Global auth check
|
||||
token = get_access_token()
|
||||
ctx = AuthContext(token=token, component=prompt)
|
||||
|
||||
if not run_auth_checks(self.auth, ctx):
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for prompt '{prompt_name}': insufficient permissions"
|
||||
|
|
|
|||
|
|
@ -1,8 +1,19 @@
|
|||
"""AggregateProvider for combining multiple providers into one.
|
||||
|
||||
This module provides `AggregateProvider` which presents multiple providers
|
||||
as a single unified provider. Used internally by FastMCP for aggregating
|
||||
components from all providers.
|
||||
This module provides `AggregateProvider`, a utility class that presents
|
||||
multiple providers as a single unified provider. Useful when you want to
|
||||
combine custom providers without creating a full FastMCP server.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from fastmcp.server.providers import AggregateProvider
|
||||
|
||||
# Combine multiple providers into one
|
||||
combined = AggregateProvider([provider1, provider2, provider3])
|
||||
|
||||
# Use like any other provider
|
||||
tools = await combined.list_tools()
|
||||
```
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -28,12 +39,15 @@ T = TypeVar("T")
|
|||
|
||||
|
||||
class AggregateProvider(Provider):
|
||||
"""Presents multiple providers as a single provider.
|
||||
"""Utility provider that combines multiple providers into one.
|
||||
|
||||
Components are aggregated from all providers. For get_* operations,
|
||||
providers are queried in parallel and the highest version is returned.
|
||||
|
||||
Errors from individual providers are logged and skipped (graceful degradation).
|
||||
|
||||
This is useful when you want to combine custom providers without creating
|
||||
a full FastMCP server.
|
||||
"""
|
||||
|
||||
def __init__(self, providers: Sequence[Provider]) -> None:
|
||||
|
|
|
|||
|
|
@ -65,13 +65,17 @@ class Provider:
|
|||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
# Visibility is the first (innermost) transform - closest to the base provider
|
||||
self._visibility = Visibility()
|
||||
self._transforms: list[Transform] = [self._visibility]
|
||||
self._transforms: list[Transform] = []
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.__class__.__name__}()"
|
||||
|
||||
@property
|
||||
def transforms(self) -> list[Transform]:
|
||||
"""All transforms including visibility (applied last/outermost)."""
|
||||
return [*self._transforms, self._visibility]
|
||||
|
||||
def add_transform(self, transform: Transform) -> None:
|
||||
"""Add a transform to this provider.
|
||||
|
||||
|
|
@ -110,7 +114,7 @@ class Provider:
|
|||
return await self.list_tools()
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.list_tools, call_next=chain)
|
||||
|
||||
return await chain()
|
||||
|
|
@ -132,7 +136,7 @@ class Provider:
|
|||
return await self.get_tool(n, version)
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_tool, call_next=chain)
|
||||
|
||||
return await chain(name, version=version)
|
||||
|
|
@ -144,7 +148,7 @@ class Provider:
|
|||
return await self.list_resources()
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.list_resources, call_next=chain)
|
||||
|
||||
return await chain()
|
||||
|
|
@ -163,7 +167,7 @@ class Provider:
|
|||
return await self.get_resource(u, version)
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_resource, call_next=chain)
|
||||
|
||||
return await chain(uri, version=version)
|
||||
|
|
@ -175,7 +179,7 @@ class Provider:
|
|||
return await self.list_resource_templates()
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.list_resource_templates, call_next=chain)
|
||||
|
||||
return await chain()
|
||||
|
|
@ -196,7 +200,7 @@ class Provider:
|
|||
return await self.get_resource_template(u, version)
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_resource_template, call_next=chain)
|
||||
|
||||
return await chain(uri, version=version)
|
||||
|
|
@ -208,7 +212,7 @@ class Provider:
|
|||
return await self.list_prompts()
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.list_prompts, call_next=chain)
|
||||
|
||||
return await chain()
|
||||
|
|
@ -227,7 +231,7 @@ class Provider:
|
|||
return await self.get_prompt(n, version)
|
||||
|
||||
chain = base
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_prompt, call_next=chain)
|
||||
|
||||
return await chain(name, version=version)
|
||||
|
|
@ -402,13 +406,13 @@ class Provider:
|
|||
async def prompts_base() -> Sequence[Prompt]:
|
||||
return prompts
|
||||
|
||||
# Apply transforms in order (first is innermost)
|
||||
# Apply transforms in order (visibility last/outermost)
|
||||
tools_chain = tools_base
|
||||
resources_chain = resources_base
|
||||
templates_chain = templates_base
|
||||
prompts_chain = prompts_base
|
||||
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
tools_chain = partial(transform.list_tools, call_next=tools_chain)
|
||||
resources_chain = partial(
|
||||
transform.list_resources, call_next=resources_chain
|
||||
|
|
|
|||
|
|
@ -19,7 +19,6 @@ from typing import TYPE_CHECKING, Any, overload
|
|||
import mcp.types
|
||||
from mcp.types import AnyUrl
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
from fastmcp.prompts.prompt import Prompt, PromptResult
|
||||
from fastmcp.resources.resource import Resource, ResourceResult
|
||||
from fastmcp.resources.template import ResourceTemplate
|
||||
|
|
@ -486,9 +485,9 @@ class FastMCPProvider(Provider):
|
|||
async def list_tools(self) -> Sequence[Tool]:
|
||||
"""List all tools from the mounted server as FastMCPProviderTools.
|
||||
|
||||
Calls the nested server's middleware to list tools, then wraps
|
||||
each tool as a FastMCPProviderTool that delegates execution to the
|
||||
nested server's middleware.
|
||||
Runs the mounted server's middleware so filtering/transformation applies.
|
||||
Wraps each tool as a FastMCPProviderTool that delegates execution to
|
||||
the nested server's middleware.
|
||||
"""
|
||||
raw_tools = await self.server.get_tools(run_middleware=True)
|
||||
return [FastMCPProviderTool.wrap(self.server, t) for t in raw_tools]
|
||||
|
|
@ -499,11 +498,11 @@ class FastMCPProvider(Provider):
|
|||
"""Get a tool by name as a FastMCPProviderTool.
|
||||
|
||||
Passes the full VersionSpec to the nested server, which handles both
|
||||
exact version matching and range filtering.
|
||||
exact version matching and range filtering. Uses _get_tool to ensure
|
||||
the nested server's transforms are applied.
|
||||
"""
|
||||
try:
|
||||
raw_tool = await self.server.get_tool(name, version)
|
||||
except NotFoundError:
|
||||
raw_tool = await self.server._get_tool(name, version)
|
||||
if raw_tool is None:
|
||||
return None
|
||||
return FastMCPProviderTool.wrap(self.server, raw_tool)
|
||||
|
||||
|
|
@ -514,9 +513,9 @@ class FastMCPProvider(Provider):
|
|||
async def list_resources(self) -> Sequence[Resource]:
|
||||
"""List all resources from the mounted server as FastMCPProviderResources.
|
||||
|
||||
Calls the nested server's middleware to list resources, then wraps
|
||||
each resource as a FastMCPProviderResource that delegates reading to the
|
||||
nested server's middleware.
|
||||
Runs the mounted server's middleware so filtering/transformation applies.
|
||||
Wraps each resource as a FastMCPProviderResource that delegates reading
|
||||
to the nested server's middleware.
|
||||
"""
|
||||
raw_resources = await self.server.get_resources(run_middleware=True)
|
||||
return [FastMCPProviderResource.wrap(self.server, r) for r in raw_resources]
|
||||
|
|
@ -527,11 +526,11 @@ class FastMCPProvider(Provider):
|
|||
"""Get a concrete resource by URI as a FastMCPProviderResource.
|
||||
|
||||
Passes the full VersionSpec to the nested server, which handles both
|
||||
exact version matching and range filtering.
|
||||
exact version matching and range filtering. Uses _get_resource to ensure
|
||||
the nested server's transforms are applied.
|
||||
"""
|
||||
try:
|
||||
raw_resource = await self.server.get_resource(uri, version)
|
||||
except NotFoundError:
|
||||
raw_resource = await self.server._get_resource(uri, version)
|
||||
if raw_resource is None:
|
||||
return None
|
||||
return FastMCPProviderResource.wrap(self.server, raw_resource)
|
||||
|
||||
|
|
@ -542,6 +541,7 @@ class FastMCPProvider(Provider):
|
|||
async def list_resource_templates(self) -> Sequence[ResourceTemplate]:
|
||||
"""List all resource templates from the mounted server.
|
||||
|
||||
Runs the mounted server's middleware so filtering/transformation applies.
|
||||
Returns FastMCPProviderResourceTemplate instances that create
|
||||
FastMCPProviderResources when materialized.
|
||||
"""
|
||||
|
|
@ -556,11 +556,11 @@ class FastMCPProvider(Provider):
|
|||
"""Get a resource template that matches the given URI.
|
||||
|
||||
Passes the full VersionSpec to the nested server, which handles both
|
||||
exact version matching and range filtering.
|
||||
exact version matching and range filtering. Uses _get_resource_template
|
||||
to ensure the nested server's transforms are applied.
|
||||
"""
|
||||
try:
|
||||
raw_template = await self.server.get_resource_template(uri, version)
|
||||
except NotFoundError:
|
||||
raw_template = await self.server._get_resource_template(uri, version)
|
||||
if raw_template is None:
|
||||
return None
|
||||
return FastMCPProviderResourceTemplate.wrap(self.server, raw_template)
|
||||
|
||||
|
|
@ -571,6 +571,7 @@ class FastMCPProvider(Provider):
|
|||
async def list_prompts(self) -> Sequence[Prompt]:
|
||||
"""List all prompts from the mounted server as FastMCPProviderPrompts.
|
||||
|
||||
Runs the mounted server's middleware so filtering/transformation applies.
|
||||
Returns FastMCPProviderPrompt instances that delegate rendering to the
|
||||
wrapped server's middleware.
|
||||
"""
|
||||
|
|
@ -583,11 +584,11 @@ class FastMCPProvider(Provider):
|
|||
"""Get a prompt by name as a FastMCPProviderPrompt.
|
||||
|
||||
Passes the full VersionSpec to the nested server, which handles both
|
||||
exact version matching and range filtering.
|
||||
exact version matching and range filtering. Uses _get_prompt to ensure
|
||||
the nested server's transforms are applied.
|
||||
"""
|
||||
try:
|
||||
raw_prompt = await self.server.get_prompt(name, version)
|
||||
except NotFoundError:
|
||||
raw_prompt = await self.server._get_prompt(name, version)
|
||||
if raw_prompt is None:
|
||||
return None
|
||||
return FastMCPProviderPrompt.wrap(self.server, raw_prompt)
|
||||
|
||||
|
|
@ -599,12 +600,12 @@ class FastMCPProvider(Provider):
|
|||
"""Return task-eligible components from the mounted server.
|
||||
|
||||
Returns the child's ACTUAL components (not wrapped) so their actual
|
||||
functions get registered with Docket. Uses _source_get_tasks() to get
|
||||
components with child server's transforms applied, then applies this
|
||||
provider's transforms for correct registration keys.
|
||||
functions get registered with Docket. Gets components with child
|
||||
server's transforms applied, then applies this provider's transforms
|
||||
for correct registration keys.
|
||||
"""
|
||||
# Get tasks with child server's transforms already applied
|
||||
components = list(await self.server._source_get_tasks())
|
||||
components = list(await self.server.get_tasks())
|
||||
|
||||
# Separate by type for this provider's transform application
|
||||
tools = [c for c in components if isinstance(c, Tool)]
|
||||
|
|
@ -631,7 +632,7 @@ class FastMCPProvider(Provider):
|
|||
templates_chain = templates_base
|
||||
prompts_chain = prompts_base
|
||||
|
||||
for transform in self._transforms:
|
||||
for transform in self.transforms:
|
||||
tools_chain = partial(transform.list_tools, call_next=tools_chain)
|
||||
resources_chain = partial(
|
||||
transform.list_resources, call_next=resources_chain
|
||||
|
|
@ -641,11 +642,16 @@ class FastMCPProvider(Provider):
|
|||
)
|
||||
prompts_chain = partial(transform.list_prompts, call_next=prompts_chain)
|
||||
|
||||
# Filter to only task-eligible components (same as base Provider)
|
||||
return [
|
||||
*await tools_chain(),
|
||||
*await resources_chain(),
|
||||
*await templates_chain(),
|
||||
*await prompts_chain(),
|
||||
c
|
||||
for c in [
|
||||
*await tools_chain(),
|
||||
*await resources_chain(),
|
||||
*await templates_chain(),
|
||||
*await prompts_chain(),
|
||||
]
|
||||
if c.task_config.supports_tasks()
|
||||
]
|
||||
|
||||
# -------------------------------------------------------------------------
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -31,6 +31,7 @@ from fastmcp.resources.template import ResourceTemplate
|
|||
from fastmcp.server.tasks.config import DEFAULT_POLL_INTERVAL_MS, DEFAULT_TTL_MS
|
||||
from fastmcp.server.tasks.keys import parse_task_key
|
||||
from fastmcp.tools.tool import Tool
|
||||
from fastmcp.utilities.versions import VersionSpec
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastmcp.server.server import FastMCP
|
||||
|
|
@ -313,16 +314,20 @@ async def tasks_result_handler(server: FastMCP, params: dict[str, Any]) -> Any:
|
|||
component: Tool | Resource | ResourceTemplate | Prompt | None = None
|
||||
try:
|
||||
if component_key.startswith("tool:"):
|
||||
name, version = _parse_key_version(component_key[5:])
|
||||
name, version_str = _parse_key_version(component_key[5:])
|
||||
version = VersionSpec(eq=version_str) if version_str else None
|
||||
component = await server.get_tool(name, version)
|
||||
elif component_key.startswith("resource:"):
|
||||
uri, version = _parse_key_version(component_key[9:])
|
||||
uri, version_str = _parse_key_version(component_key[9:])
|
||||
version = VersionSpec(eq=version_str) if version_str else None
|
||||
component = await server.get_resource(uri, version)
|
||||
elif component_key.startswith("template:"):
|
||||
uri, version = _parse_key_version(component_key[9:])
|
||||
uri, version_str = _parse_key_version(component_key[9:])
|
||||
version = VersionSpec(eq=version_str) if version_str else None
|
||||
component = await server.get_resource_template(uri, version)
|
||||
elif component_key.startswith("prompt:"):
|
||||
name, version = _parse_key_version(component_key[7:])
|
||||
name, version_str = _parse_key_version(component_key[7:])
|
||||
version = VersionSpec(eq=version_str) if version_str else None
|
||||
component = await server.get_prompt(name, version)
|
||||
except NotFoundError:
|
||||
component = None
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ T = TypeVar("T", default=Any)
|
|||
class FastMCPMeta(TypedDict, total=False):
|
||||
tags: list[str]
|
||||
version: str
|
||||
versions: list[str]
|
||||
|
||||
|
||||
def get_fastmcp_metadata(meta: dict[str, Any] | None) -> FastMCPMeta:
|
||||
|
|
|
|||
|
|
@ -106,11 +106,11 @@ async def inspect_fastmcp_v2(mcp: FastMCP[Any]) -> FastMCPInfo:
|
|||
Returns:
|
||||
FastMCPInfo dataclass containing the extracted information
|
||||
"""
|
||||
# Get all components directly without middleware (auth, rate limiting, etc.)
|
||||
tools_list = await mcp.get_tools(run_middleware=False)
|
||||
prompts_list = await mcp.get_prompts(run_middleware=False)
|
||||
resources_list = await mcp.get_resources(run_middleware=False)
|
||||
templates_list = await mcp.get_resource_templates(run_middleware=False)
|
||||
# Get all components
|
||||
tools_list = await mcp.get_tools()
|
||||
prompts_list = await mcp.get_prompts()
|
||||
resources_list = await mcp.get_resources()
|
||||
templates_list = await mcp.get_resource_templates()
|
||||
|
||||
# Extract detailed tool information
|
||||
tool_infos = []
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
|
|||
|
||||
from fastmcp import FastMCP
|
||||
from fastmcp.client import Client
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
from fastmcp.server.auth import (
|
||||
AccessToken,
|
||||
AuthContext,
|
||||
|
|
@ -302,15 +301,17 @@ class TestToolLevelAuth:
|
|||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
||||
async def test_get_tool_returns_not_found_without_auth(self):
|
||||
async def test_get_tool_returns_none_without_auth(self):
|
||||
"""get_tool() returns None for unauthorized tools (consistent with list filtering)."""
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool(auth=require_auth)
|
||||
def protected_tool() -> str:
|
||||
return "protected"
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.get_tool("protected_tool")
|
||||
# get_tool() returns None for unauthorized tools
|
||||
tool = await mcp.get_tool("protected_tool")
|
||||
assert tool is None
|
||||
|
||||
async def test_get_tool_returns_tool_with_auth(self):
|
||||
mcp = FastMCP()
|
||||
|
|
@ -323,6 +324,7 @@ class TestToolLevelAuth:
|
|||
tok = set_token(token)
|
||||
try:
|
||||
tool = await mcp.get_tool("protected_tool")
|
||||
assert tool is not None
|
||||
assert tool.name == "protected_tool"
|
||||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
|
@ -334,6 +336,12 @@ class TestToolLevelAuth:
|
|||
|
||||
|
||||
class TestAuthMiddleware:
|
||||
"""Tests for middleware filtering via MCP handler layer.
|
||||
|
||||
These tests call _list_tools_mcp() which applies middleware during list,
|
||||
simulating what happens when a client calls list_tools over MCP.
|
||||
"""
|
||||
|
||||
async def test_middleware_filters_tools_without_token(self):
|
||||
mcp = FastMCP(middleware=[AuthMiddleware(auth=require_auth)])
|
||||
|
||||
|
|
@ -342,7 +350,7 @@ class TestAuthMiddleware:
|
|||
return "public"
|
||||
|
||||
# No token - all tools filtered by middleware
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 0
|
||||
|
||||
async def test_middleware_allows_tools_with_token(self):
|
||||
|
|
@ -355,7 +363,7 @@ class TestAuthMiddleware:
|
|||
token = make_token()
|
||||
tok = set_token(token)
|
||||
try:
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 1
|
||||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
|
@ -371,7 +379,7 @@ class TestAuthMiddleware:
|
|||
token = make_token(scopes=["read"])
|
||||
tok = set_token(token)
|
||||
try:
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 0
|
||||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
|
@ -380,7 +388,7 @@ class TestAuthMiddleware:
|
|||
token = make_token(scopes=["api"])
|
||||
tok = set_token(token)
|
||||
try:
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 1
|
||||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
|
@ -399,7 +407,7 @@ class TestAuthMiddleware:
|
|||
return "admin"
|
||||
|
||||
# No token - public tool allowed, admin tool blocked
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 1
|
||||
assert tools[0].name == "public_tool"
|
||||
|
||||
|
|
@ -407,7 +415,7 @@ class TestAuthMiddleware:
|
|||
token = make_token(scopes=["admin"])
|
||||
tok = set_token(token)
|
||||
try:
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp._list_tools_mcp()
|
||||
assert len(tools) == 2
|
||||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import pytest
|
|||
from mcp.types import TextContent
|
||||
|
||||
from fastmcp import Client, Context, FastMCP
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
from fastmcp.prompts.prompt import Prompt
|
||||
|
||||
|
||||
|
|
@ -352,9 +351,6 @@ class TestPromptEnabled:
|
|||
prompts = await mcp.get_prompts()
|
||||
assert len(prompts) == 0
|
||||
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("sample_prompt")
|
||||
|
||||
async def test_prompt_toggle_enabled(self):
|
||||
mcp = FastMCP()
|
||||
|
||||
|
|
@ -381,8 +377,9 @@ class TestPromptEnabled:
|
|||
prompts = await mcp.get_prompts()
|
||||
assert len(prompts) == 0
|
||||
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("sample_prompt")
|
||||
# _get_prompt() applies visibility transform, returns None for disabled
|
||||
prompt = await mcp._get_prompt("sample_prompt")
|
||||
assert prompt is None
|
||||
|
||||
async def test_get_prompt_and_disable(self):
|
||||
mcp = FastMCP()
|
||||
|
|
@ -391,15 +388,16 @@ class TestPromptEnabled:
|
|||
def sample_prompt() -> str:
|
||||
return "Hello, world!"
|
||||
|
||||
prompt = await mcp.get_prompt("sample_prompt")
|
||||
prompt = await mcp._get_prompt("sample_prompt")
|
||||
assert prompt is not None
|
||||
|
||||
mcp.disable(keys=["prompt:sample_prompt@"])
|
||||
prompts = await mcp.get_prompts()
|
||||
assert len(prompts) == 0
|
||||
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("sample_prompt")
|
||||
# _get_prompt() applies visibility transform, returns None for disabled
|
||||
prompt = await mcp._get_prompt("sample_prompt")
|
||||
assert prompt is None
|
||||
|
||||
async def test_cant_get_disabled_prompt(self):
|
||||
mcp = FastMCP()
|
||||
|
|
@ -410,8 +408,9 @@ class TestPromptEnabled:
|
|||
|
||||
mcp.disable(keys=["prompt:sample_prompt@"])
|
||||
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("sample_prompt")
|
||||
# _get_prompt() applies visibility transform, returns None for disabled
|
||||
prompt = await mcp._get_prompt("sample_prompt")
|
||||
assert prompt is None
|
||||
|
||||
|
||||
class TestPromptTags:
|
||||
|
|
@ -455,18 +454,20 @@ class TestPromptTags:
|
|||
|
||||
async def test_read_prompt_includes_tags(self):
|
||||
mcp = self.create_server(include_tags={"a"})
|
||||
prompt = await mcp.get_prompt("prompt_1")
|
||||
# _get_prompt applies visibility transform (tag filtering)
|
||||
prompt = await mcp._get_prompt("prompt_1")
|
||||
result = await prompt.render({})
|
||||
assert result.messages[0].content.text == "1"
|
||||
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("prompt_2")
|
||||
prompt = await mcp._get_prompt("prompt_2")
|
||||
assert prompt is None
|
||||
|
||||
async def test_read_prompt_excludes_tags(self):
|
||||
mcp = self.create_server(exclude_tags={"a"})
|
||||
with pytest.raises(NotFoundError, match="Unknown prompt"):
|
||||
await mcp.get_prompt("prompt_1")
|
||||
# _get_prompt applies visibility transform (tag filtering)
|
||||
prompt = await mcp._get_prompt("prompt_1")
|
||||
assert prompt is None
|
||||
|
||||
prompt = await mcp.get_prompt("prompt_2")
|
||||
prompt = await mcp._get_prompt("prompt_2")
|
||||
result = await prompt.render({})
|
||||
assert result.messages[0].content.text == "2"
|
||||
|
|
|
|||
|
|
@ -838,8 +838,8 @@ class TestAsProxyKwarg:
|
|||
provider = mcp._providers[1]
|
||||
# With namespace, we get FastMCPProvider with a Namespace layer
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub
|
||||
|
||||
async def test_as_proxy_false(self):
|
||||
|
|
@ -852,8 +852,8 @@ class TestAsProxyKwarg:
|
|||
provider = mcp._providers[1]
|
||||
# With namespace, we get FastMCPProvider with a Namespace layer
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub
|
||||
|
||||
async def test_as_proxy_true(self):
|
||||
|
|
@ -866,8 +866,8 @@ class TestAsProxyKwarg:
|
|||
provider = mcp._providers[1]
|
||||
# With namespace, we get FastMCPProvider with a Namespace layer
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is not sub
|
||||
assert isinstance(provider.server, FastMCPProxy)
|
||||
|
||||
|
|
@ -891,8 +891,8 @@ class TestAsProxyKwarg:
|
|||
# Index 1 because LocalProvider is at index 0
|
||||
provider = mcp._providers[1]
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub
|
||||
|
||||
async def test_as_proxy_ignored_for_proxy_mounts_default(self):
|
||||
|
|
@ -905,8 +905,8 @@ class TestAsProxyKwarg:
|
|||
# Index 1 because LocalProvider is at index 0
|
||||
provider = mcp._providers[1]
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub_proxy
|
||||
|
||||
async def test_as_proxy_ignored_for_proxy_mounts_false(self):
|
||||
|
|
@ -919,8 +919,8 @@ class TestAsProxyKwarg:
|
|||
# Index 1 because LocalProvider is at index 0
|
||||
provider = mcp._providers[1]
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub_proxy
|
||||
|
||||
async def test_as_proxy_ignored_for_proxy_mounts_true(self):
|
||||
|
|
@ -933,8 +933,8 @@ class TestAsProxyKwarg:
|
|||
# Index 1 because LocalProvider is at index 0
|
||||
provider = mcp._providers[1]
|
||||
assert isinstance(provider, FastMCPProvider)
|
||||
assert len(provider._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider._transforms[1], Namespace)
|
||||
assert len(provider._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider._transforms[0], Namespace)
|
||||
assert provider.server is sub_proxy
|
||||
|
||||
async def test_as_proxy_mounts_still_have_live_link(self):
|
||||
|
|
@ -1169,20 +1169,20 @@ class TestCustomRouteForwarding:
|
|||
# LocalProvider is at index 0, mounted provider at index 1
|
||||
provider1 = main_server._providers[1]
|
||||
assert isinstance(provider1, FastMCPProvider)
|
||||
assert len(provider1._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider1._transforms[1], Namespace)
|
||||
assert len(provider1._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider1._transforms[0], Namespace)
|
||||
assert provider1.server == sub_server1
|
||||
assert provider1._transforms[1]._prefix == "sub1"
|
||||
assert provider1._transforms[0]._prefix == "sub1"
|
||||
|
||||
# Mount second server
|
||||
main_server.mount(sub_server2, "sub2")
|
||||
assert len(main_server._providers) == 3
|
||||
provider2 = main_server._providers[2]
|
||||
assert isinstance(provider2, FastMCPProvider)
|
||||
assert len(provider2._transforms) == 2 # Visibility + Namespace
|
||||
assert isinstance(provider2._transforms[1], Namespace)
|
||||
assert len(provider2._transforms) == 1 # Just Namespace
|
||||
assert isinstance(provider2._transforms[0], Namespace)
|
||||
assert provider2.server == sub_server2
|
||||
assert provider2._transforms[1]._prefix == "sub2"
|
||||
assert provider2._transforms[0]._prefix == "sub2"
|
||||
|
||||
async def test_multiple_routes_same_server(self):
|
||||
"""Test that multiple custom routes from same server are all included."""
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ async def test_transformed_tool_filtering():
|
|||
# Enable only tools with the enabled_tools tag
|
||||
mcp.enable(tags={"enabled_tools"}, only=True)
|
||||
|
||||
tools = await mcp.get_tools(run_middleware=True)
|
||||
tools = await mcp.get_tools()
|
||||
# With transformation applied, the tool now has the enabled_tools tag
|
||||
assert len(tools) == 1
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from mcp.types import TextContent
|
|||
from fastmcp import FastMCP
|
||||
from fastmcp.utilities.versions import (
|
||||
VersionKey,
|
||||
VersionSpec,
|
||||
compare_versions,
|
||||
is_version_greater,
|
||||
)
|
||||
|
|
@ -377,7 +378,7 @@ class TestMountedServerVersioning:
|
|||
assert tool.version == "2.0"
|
||||
|
||||
# Get specific version
|
||||
tool_v1 = await parent.get_tool("child_calc", version="1.0")
|
||||
tool_v1 = await parent.get_tool("child_calc", VersionSpec(eq="1.0"))
|
||||
assert tool_v1 is not None
|
||||
assert tool_v1.version == "1.0"
|
||||
|
||||
|
|
@ -482,18 +483,13 @@ class TestVersionFilter:
|
|||
assert len(tools) == 1
|
||||
assert tools[0].version == "3.0"
|
||||
|
||||
# Can request specific versions in range
|
||||
tool_v2 = await mcp.get_tool("add", version="2.0")
|
||||
# Can request specific versions in range (use _get_tool to apply transforms)
|
||||
tool_v2 = await mcp._get_tool("add", VersionSpec(eq="2.0"))
|
||||
assert tool_v2 is not None
|
||||
assert tool_v2.version == "2.0"
|
||||
|
||||
# Cannot request version outside range
|
||||
import pytest
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.get_tool("add", version="1.0")
|
||||
# Cannot request version outside range - returns None
|
||||
assert await mcp._get_tool("add", VersionSpec(eq="1.0")) is None
|
||||
|
||||
async def test_version_range(self):
|
||||
"""VersionFilter(version_gte='2.0', version_lt='3.0') shows only v2.x."""
|
||||
|
|
@ -524,21 +520,14 @@ class TestVersionFilter:
|
|||
assert len(tools) == 1
|
||||
assert tools[0].version == "2.5"
|
||||
|
||||
# Can request specific versions in range
|
||||
tool_v2 = await mcp.get_tool("calc", version="2.0")
|
||||
# Can request specific versions in range (use _get_tool to apply transforms)
|
||||
tool_v2 = await mcp._get_tool("calc", VersionSpec(eq="2.0"))
|
||||
assert tool_v2 is not None
|
||||
assert tool_v2.version == "2.0"
|
||||
|
||||
# Versions outside range are not accessible
|
||||
import pytest
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.get_tool("calc", version="1.0")
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.get_tool("calc", version="3.0")
|
||||
# Versions outside range are not accessible - return None
|
||||
assert await mcp._get_tool("calc", VersionSpec(eq="1.0")) is None
|
||||
assert await mcp._get_tool("calc", VersionSpec(eq="3.0")) is None
|
||||
|
||||
async def test_unversioned_always_passes(self):
|
||||
"""Unversioned components pass through any filter."""
|
||||
|
|
@ -588,10 +577,8 @@ class TestVersionFilter:
|
|||
assert tools[0].version == "2025-01-01"
|
||||
|
||||
async def test_get_tool_respects_filter(self):
|
||||
"""get_tool() raises NotFoundError if highest version is filtered out."""
|
||||
import pytest
|
||||
"""get_tool() returns None if highest version is filtered out."""
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
from fastmcp.server.transforms import VersionFilter
|
||||
|
||||
mcp = FastMCP()
|
||||
|
|
@ -602,9 +589,8 @@ class TestVersionFilter:
|
|||
|
||||
mcp.add_transform(VersionFilter(version_lt="3.0"))
|
||||
|
||||
# Tool exists but is filtered out
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.get_tool("only_v5")
|
||||
# Tool exists but is filtered out - returns None (use _get_tool to apply transforms)
|
||||
assert await mcp._get_tool("only_v5") is None
|
||||
|
||||
async def test_must_specify_at_least_one(self):
|
||||
"""VersionFilter() with no args raises ValueError."""
|
||||
|
|
@ -845,13 +831,8 @@ class TestMountedVersionFiltering:
|
|||
tools = await parent.get_tools()
|
||||
assert len(tools) == 0
|
||||
|
||||
# get_tool should also return None (respects filter)
|
||||
import pytest
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await parent.get_tool("child_high_version_tool")
|
||||
# _get_tool should also return None (respects filter, applies transforms)
|
||||
assert await parent._get_tool("child_high_version_tool") is None
|
||||
|
||||
|
||||
class TestMountedRangeFiltering:
|
||||
|
|
@ -876,7 +857,8 @@ class TestMountedRangeFiltering:
|
|||
parent.add_transform(VersionFilter(version_lt="2.0"))
|
||||
|
||||
# Should return v1.0 (the highest version that matches <2.0)
|
||||
tool = await parent.get_tool("child_calc")
|
||||
# Use _get_tool to apply transforms
|
||||
tool = await parent._get_tool("child_calc")
|
||||
assert tool is not None
|
||||
assert tool.version == "1.0"
|
||||
|
||||
|
|
@ -902,18 +884,14 @@ class TestMountedRangeFiltering:
|
|||
parent.mount(child, "child")
|
||||
parent.add_transform(VersionFilter(version_gte="1.0", version_lt="3.0"))
|
||||
|
||||
# Request specific version within range
|
||||
tool = await parent.get_tool("child_calc", version="1.0")
|
||||
# Request specific version within range (use _get_tool to apply transforms)
|
||||
tool = await parent._get_tool("child_calc", VersionSpec(eq="1.0"))
|
||||
assert tool is not None
|
||||
assert tool.version == "1.0"
|
||||
|
||||
# Request version outside range should fail
|
||||
import pytest
|
||||
|
||||
from fastmcp.exceptions import NotFoundError
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await parent.get_tool("child_calc", version="3.0")
|
||||
# Request version outside range should return None
|
||||
result = await parent._get_tool("child_calc", VersionSpec(eq="3.0"))
|
||||
assert result is None
|
||||
|
||||
|
||||
class TestUnversionedExemption:
|
||||
|
|
@ -954,7 +932,7 @@ class TestUnversionedExemption:
|
|||
|
||||
# Even with explicit version request, unversioned tool is returned
|
||||
# (it's the only version that exists, and unversioned matches any spec)
|
||||
tool = await mcp.get_tool("my_tool", version="1.0")
|
||||
tool = await mcp.get_tool("my_tool", VersionSpec(eq="1.0"))
|
||||
assert tool is not None
|
||||
assert tool.version is None
|
||||
|
||||
|
|
@ -1090,12 +1068,16 @@ class TestVersionedCalls:
|
|||
assert result.structured_content["result"] == 12
|
||||
|
||||
# Explicit v1.0 (addition)
|
||||
result = await mcp.call_tool("calculate", {"x": 3, "y": 4}, version="1.0")
|
||||
result = await mcp.call_tool(
|
||||
"calculate", {"x": 3, "y": 4}, version=VersionSpec(eq="1.0")
|
||||
)
|
||||
assert result.structured_content is not None
|
||||
assert result.structured_content["result"] == 7
|
||||
|
||||
# Explicit v2.0 (multiplication)
|
||||
result = await mcp.call_tool("calculate", {"x": 3, "y": 4}, version="2.0")
|
||||
result = await mcp.call_tool(
|
||||
"calculate", {"x": 3, "y": 4}, version=VersionSpec(eq="2.0")
|
||||
)
|
||||
assert result.structured_content is not None
|
||||
assert result.structured_content["result"] == 12
|
||||
|
||||
|
|
@ -1116,7 +1098,7 @@ class TestVersionedCalls:
|
|||
assert result.contents[0].content == "config v2"
|
||||
|
||||
# Explicit v1.0
|
||||
result = await mcp.read_resource("data://config", version="1.0")
|
||||
result = await mcp.read_resource("data://config", version=VersionSpec(eq="1.0"))
|
||||
assert result.contents[0].content == "config v1"
|
||||
|
||||
async def test_render_prompt_with_version(self):
|
||||
|
|
@ -1137,7 +1119,7 @@ class TestVersionedCalls:
|
|||
assert isinstance(content, TextContent) and content.text == "Hello from v2"
|
||||
|
||||
# Explicit v1.0
|
||||
result = await mcp.render_prompt("greet", version="1.0")
|
||||
result = await mcp.render_prompt("greet", version=VersionSpec(eq="1.0"))
|
||||
content = result.messages[0].content
|
||||
assert isinstance(content, TextContent) and content.text == "Hello from v1"
|
||||
|
||||
|
|
@ -1154,7 +1136,7 @@ class TestVersionedCalls:
|
|||
return "v1"
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
await mcp.call_tool("mytool", {}, version="999.0")
|
||||
await mcp.call_tool("mytool", {}, version=VersionSpec(eq="999.0"))
|
||||
|
||||
|
||||
class TestClientVersionSelection:
|
||||
|
|
|
|||
|
|
@ -806,6 +806,7 @@ class TestProxy:
|
|||
# when adding transformed tools to proxy servers. Needs separate investigation.
|
||||
|
||||
add_tool = await proxy_server.get_tool("add")
|
||||
assert add_tool is not None
|
||||
new_add_tool = Tool.from_tool(
|
||||
add_tool,
|
||||
name="add_transformed",
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue