mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-20 20:44:17 +02:00
Refactor FastMCP to use inherited _get_* methods from Provider
- get_*() now does aggregation + component auth (raises AuthorizationError) - Deleted _get_*() overrides - inherited from Provider applies transforms - Simplified AuthMiddleware to global auth only - Changed version params to VersionSpec | None (not str | None) - Updated tests to use _get_*() where visibility filtering is expected
This commit is contained in:
parent
5ade5a4838
commit
5b24ea393d
7 changed files with 222 additions and 298 deletions
|
|
@ -136,16 +136,16 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for tool '{tool_name}': missing context"
|
||||
)
|
||||
|
||||
# Resolve tool (includes component-level auth check)
|
||||
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"
|
||||
|
|
@ -201,18 +201,18 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for resource '{uri}': missing context"
|
||||
)
|
||||
|
||||
# Try concrete resource first, then template (includes component-level auth check)
|
||||
component = await fastmcp.fastmcp.get_resource(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))
|
||||
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"
|
||||
|
|
@ -294,16 +294,16 @@ class AuthMiddleware(Middleware):
|
|||
f"Authorization failed for prompt '{prompt_name}': missing context"
|
||||
)
|
||||
|
||||
# Resolve prompt (includes component-level auth check)
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -628,7 +628,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
|
||||
|
|
|
|||
|
|
@ -356,7 +356,6 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
)
|
||||
|
||||
# Local provider is always first in the provider list
|
||||
# Note: _transforms is initialized by Provider.__init__()
|
||||
self._providers: list[Provider] = [
|
||||
self._local_provider,
|
||||
*(providers or []),
|
||||
|
|
@ -1127,78 +1126,58 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
return _dedupe_with_versions(authorized, lambda t: t.name)
|
||||
|
||||
async def get_tool(
|
||||
self, name: str, version: VersionSpec | str | None = None
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Tool | None:
|
||||
"""Get a tool by name with all server transforms applied.
|
||||
"""Get a tool by name via aggregation from providers.
|
||||
|
||||
Returns None if not found or if the tool is disabled via visibility settings.
|
||||
This is the raw lookup that Provider._get_tool() wraps with transforms.
|
||||
Aggregates from all sub-providers and applies component-level auth.
|
||||
|
||||
Args:
|
||||
name: The tool name.
|
||||
version: Version filter. Can be:
|
||||
- None: returns highest version
|
||||
- str: returns exact version match
|
||||
- VersionSpec: returns best match within spec (highest matching)
|
||||
"""
|
||||
# Convert string to VersionSpec for backward compatibility
|
||||
if isinstance(version, str):
|
||||
version_spec: VersionSpec | None = VersionSpec(eq=version)
|
||||
else:
|
||||
version_spec = version
|
||||
version: Version filter (None returns highest version).
|
||||
|
||||
# Use _get_tool which applies server transforms
|
||||
tool = await self._get_tool(name, version_spec)
|
||||
if tool is None:
|
||||
Returns:
|
||||
The tool if found and authorized, None if not found.
|
||||
|
||||
Raises:
|
||||
AuthorizationError: If component-level auth fails.
|
||||
"""
|
||||
|
||||
# Aggregate from all sub-providers (each applies their own transforms)
|
||||
results = await gather(
|
||||
*[p._get_tool(name, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Tool] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_tool({name!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
# Auth check
|
||||
tool: Tool = max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Component auth - raises if unauthorized
|
||||
skip_auth, token = _get_auth_context()
|
||||
if not skip_auth and tool.auth is not None:
|
||||
ctx = AuthContext(token=token, component=tool)
|
||||
if not run_auth_checks(tool.auth, ctx):
|
||||
return None
|
||||
raise AuthorizationError(f"Unauthorized access to tool: {name!r}")
|
||||
|
||||
return tool
|
||||
|
||||
async def _get_tool(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Tool | None:
|
||||
"""Get tool with all transforms applied (server transforms + aggregation).
|
||||
|
||||
Overrides Provider._get_tool to aggregate from providers after applying
|
||||
server-level transforms.
|
||||
"""
|
||||
|
||||
async def base(n: str, *, version: VersionSpec | None = None) -> Tool | None:
|
||||
# Aggregate from all sub-providers
|
||||
results = await gather(
|
||||
*[p._get_tool(n, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Tool] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_tool({n!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
return max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Build transform chain: server transforms applied over aggregation
|
||||
chain = base
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_tool, call_next=chain)
|
||||
|
||||
return await chain(name, version=version)
|
||||
# _get_tool is inherited from Provider - wraps get_tool() with transforms
|
||||
|
||||
async def get_resources(self, *, run_middleware: bool = False) -> list[Resource]:
|
||||
"""Get all enabled resources from providers.
|
||||
|
|
@ -1245,80 +1224,57 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
return _dedupe_with_versions(authorized, lambda r: str(r.uri))
|
||||
|
||||
async def get_resource(
|
||||
self, uri: str, version: VersionSpec | str | None = None
|
||||
self, uri: str, version: VersionSpec | None = None
|
||||
) -> Resource | None:
|
||||
"""Get a resource by URI with all server transforms applied.
|
||||
"""Get a resource by URI via aggregation from providers.
|
||||
|
||||
Returns None if not found or if the resource is disabled via visibility settings.
|
||||
This is the raw lookup that Provider._get_resource() wraps with transforms.
|
||||
Aggregates from all sub-providers and applies component-level auth.
|
||||
|
||||
Args:
|
||||
uri: The resource URI.
|
||||
version: Version filter. Can be:
|
||||
- None: returns highest version
|
||||
- str: returns exact version match
|
||||
- VersionSpec: returns best match within spec (highest matching)
|
||||
"""
|
||||
# Convert string to VersionSpec for backward compatibility
|
||||
if isinstance(version, str):
|
||||
version_spec: VersionSpec | None = VersionSpec(eq=version)
|
||||
else:
|
||||
version_spec = version
|
||||
version: Version filter (None returns highest version).
|
||||
|
||||
# Use _get_resource which applies server transforms
|
||||
resource = await self._get_resource(uri, version_spec)
|
||||
if resource is None:
|
||||
Returns:
|
||||
The resource if found and authorized, None if not found.
|
||||
|
||||
Raises:
|
||||
AuthorizationError: If component-level auth fails.
|
||||
"""
|
||||
# Aggregate from all sub-providers (each applies their own transforms)
|
||||
results = await gather(
|
||||
*[p._get_resource(uri, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Resource] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_resource({uri!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
# Auth check
|
||||
resource: Resource = max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Component auth - raises if unauthorized
|
||||
skip_auth, token = _get_auth_context()
|
||||
if not skip_auth and resource.auth is not None:
|
||||
ctx = AuthContext(token=token, component=resource)
|
||||
if not run_auth_checks(resource.auth, ctx):
|
||||
return None
|
||||
raise AuthorizationError(f"Unauthorized access to resource: {uri!r}")
|
||||
|
||||
return resource
|
||||
|
||||
async def _get_resource(
|
||||
self, uri: str, version: VersionSpec | None = None
|
||||
) -> Resource | None:
|
||||
"""Get resource with all transforms applied (server transforms + aggregation).
|
||||
|
||||
Overrides Provider._get_resource to aggregate from providers after applying
|
||||
server-level transforms.
|
||||
"""
|
||||
|
||||
async def base(
|
||||
u: str, *, version: VersionSpec | None = None
|
||||
) -> Resource | None:
|
||||
# Aggregate from all sub-providers
|
||||
results = await gather(
|
||||
*[p._get_resource(u, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Resource] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_resource({u!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
return max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Build transform chain: server transforms applied over aggregation
|
||||
chain = base
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_resource, call_next=chain)
|
||||
|
||||
return await chain(uri, version=version)
|
||||
# _get_resource is inherited from Provider - wraps get_resource() with transforms
|
||||
|
||||
async def get_resource_templates(
|
||||
self, *, run_middleware: bool = False
|
||||
|
|
@ -1369,80 +1325,57 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
return _dedupe_with_versions(authorized, lambda t: t.uri_template)
|
||||
|
||||
async def get_resource_template(
|
||||
self, uri: str, version: VersionSpec | str | None = None
|
||||
self, uri: str, version: VersionSpec | None = None
|
||||
) -> ResourceTemplate | None:
|
||||
"""Get a resource template by URI with all server transforms applied.
|
||||
"""Get a resource template by URI via aggregation from providers.
|
||||
|
||||
Returns None if not found or if the template is disabled via visibility settings.
|
||||
This is the raw lookup that Provider._get_resource_template() wraps with transforms.
|
||||
Aggregates from all sub-providers and applies component-level auth.
|
||||
|
||||
Args:
|
||||
uri: The template URI to match.
|
||||
version: Version filter. Can be:
|
||||
- None: returns highest version
|
||||
- str: returns exact version match
|
||||
- VersionSpec: returns best match within spec (highest matching)
|
||||
"""
|
||||
# Convert string to VersionSpec for backward compatibility
|
||||
if isinstance(version, str):
|
||||
version_spec: VersionSpec | None = VersionSpec(eq=version)
|
||||
else:
|
||||
version_spec = version
|
||||
version: Version filter (None returns highest version).
|
||||
|
||||
# Use _get_resource_template which applies server transforms
|
||||
template = await self._get_resource_template(uri, version_spec)
|
||||
if template is None:
|
||||
Returns:
|
||||
The template if found and authorized, None if not found.
|
||||
|
||||
Raises:
|
||||
AuthorizationError: If component-level auth fails.
|
||||
"""
|
||||
# Aggregate from all sub-providers (each applies their own transforms)
|
||||
results = await gather(
|
||||
*[p._get_resource_template(uri, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[ResourceTemplate] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_resource_template({uri!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
# Auth check
|
||||
template: ResourceTemplate = max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Component auth - raises if unauthorized
|
||||
skip_auth, token = _get_auth_context()
|
||||
if not skip_auth and template.auth is not None:
|
||||
ctx = AuthContext(token=token, component=template)
|
||||
if not run_auth_checks(template.auth, ctx):
|
||||
return None
|
||||
raise AuthorizationError(f"Unauthorized access to template: {uri!r}")
|
||||
|
||||
return template
|
||||
|
||||
async def _get_resource_template(
|
||||
self, uri: str, version: VersionSpec | None = None
|
||||
) -> ResourceTemplate | None:
|
||||
"""Get resource template with all transforms applied.
|
||||
|
||||
Overrides Provider._get_resource_template to aggregate from providers after
|
||||
applying server-level transforms.
|
||||
"""
|
||||
|
||||
async def base(
|
||||
u: str, *, version: VersionSpec | None = None
|
||||
) -> ResourceTemplate | None:
|
||||
# Aggregate from all sub-providers
|
||||
results = await gather(
|
||||
*[p._get_resource_template(u, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[ResourceTemplate] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_resource_template({u!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
return max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Build transform chain: server transforms applied over aggregation
|
||||
chain = base
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_resource_template, call_next=chain)
|
||||
|
||||
return await chain(uri, version=version)
|
||||
# _get_resource_template is inherited from Provider - wraps get_resource_template() with transforms
|
||||
|
||||
async def get_prompts(self, *, run_middleware: bool = False) -> list[Prompt]:
|
||||
"""Get all enabled prompts from providers.
|
||||
|
|
@ -1489,78 +1422,57 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
return _dedupe_with_versions(authorized, lambda p: p.name)
|
||||
|
||||
async def get_prompt(
|
||||
self, name: str, version: VersionSpec | str | None = None
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a prompt by name with all server transforms applied.
|
||||
"""Get a prompt by name via aggregation from providers.
|
||||
|
||||
Returns None if not found or if the prompt is disabled via visibility settings.
|
||||
This is the raw lookup that Provider._get_prompt() wraps with transforms.
|
||||
Aggregates from all sub-providers and applies component-level auth.
|
||||
|
||||
Args:
|
||||
name: The prompt name.
|
||||
version: Version filter. Can be:
|
||||
- None: returns highest version
|
||||
- str: returns exact version match
|
||||
- VersionSpec: returns best match within spec (highest matching)
|
||||
"""
|
||||
# Convert string to VersionSpec for backward compatibility
|
||||
if isinstance(version, str):
|
||||
version_spec: VersionSpec | None = VersionSpec(eq=version)
|
||||
else:
|
||||
version_spec = version
|
||||
version: Version filter (None returns highest version).
|
||||
|
||||
# Use _get_prompt which applies server transforms
|
||||
prompt = await self._get_prompt(name, version_spec)
|
||||
if prompt is None:
|
||||
Returns:
|
||||
The prompt if found and authorized, None if not found.
|
||||
|
||||
Raises:
|
||||
AuthorizationError: If component-level auth fails.
|
||||
"""
|
||||
# Aggregate from all sub-providers (each applies their own transforms)
|
||||
results = await gather(
|
||||
*[p._get_prompt(name, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Prompt] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_prompt({name!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
# Auth check
|
||||
prompt: Prompt = max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Component auth - raises if unauthorized
|
||||
skip_auth, token = _get_auth_context()
|
||||
if not skip_auth and prompt.auth is not None:
|
||||
ctx = AuthContext(token=token, component=prompt)
|
||||
if not run_auth_checks(prompt.auth, ctx):
|
||||
return None
|
||||
raise AuthorizationError(f"Unauthorized access to prompt: {name!r}")
|
||||
|
||||
return prompt
|
||||
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get prompt with all transforms applied (server transforms + aggregation).
|
||||
|
||||
Overrides Provider._get_prompt to aggregate from providers after applying
|
||||
server-level transforms.
|
||||
"""
|
||||
|
||||
async def base(n: str, *, version: VersionSpec | None = None) -> Prompt | None:
|
||||
# Aggregate from all sub-providers
|
||||
results = await gather(
|
||||
*[p._get_prompt(n, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
# Collect valid results, pick highest version
|
||||
valid: list[Prompt] = []
|
||||
for i, result in enumerate(results):
|
||||
if isinstance(result, BaseException):
|
||||
if not isinstance(result, NotFoundError):
|
||||
logger.debug(
|
||||
f"Error during get_prompt({n!r}) from provider "
|
||||
f"{self._providers[i]}: {result}"
|
||||
)
|
||||
continue
|
||||
if result is not None:
|
||||
valid.append(result)
|
||||
|
||||
if not valid:
|
||||
return None
|
||||
|
||||
return max(valid, key=version_sort_key) # type: ignore[type-var]
|
||||
|
||||
# Build transform chain: server transforms applied over aggregation
|
||||
chain = base
|
||||
for transform in self.transforms:
|
||||
chain = partial(transform.get_prompt, call_next=chain)
|
||||
|
||||
return await chain(name, version=version)
|
||||
# _get_prompt is inherited from Provider - wraps get_prompt() with transforms
|
||||
|
||||
@overload
|
||||
async def call_tool(
|
||||
|
|
@ -1568,7 +1480,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: None = None,
|
||||
) -> ToolResult: ...
|
||||
|
|
@ -1579,7 +1491,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta,
|
||||
) -> mcp.types.CreateTaskResult: ...
|
||||
|
|
@ -1589,7 +1501,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta | None = None,
|
||||
) -> ToolResult | mcp.types.CreateTaskResult:
|
||||
|
|
@ -1643,10 +1555,11 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
)
|
||||
|
||||
# Core logic: find and execute tool (providers queried in parallel)
|
||||
# Use _get_tool to apply transforms (including visibility)
|
||||
with server_span(
|
||||
f"tools/call {name}", "tools/call", self.name, "tool", name
|
||||
) as span:
|
||||
tool = await self.get_tool(name, version=version)
|
||||
tool = await self._get_tool(name, version=version)
|
||||
if tool is None:
|
||||
raise NotFoundError(f"Unknown tool: {name!r}")
|
||||
span.set_attributes(tool.get_span_attributes())
|
||||
|
|
@ -1671,7 +1584,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
self,
|
||||
uri: str,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: None = None,
|
||||
) -> ResourceResult: ...
|
||||
|
|
@ -1681,7 +1594,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
self,
|
||||
uri: str,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta,
|
||||
) -> mcp.types.CreateTaskResult: ...
|
||||
|
|
@ -1690,7 +1603,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
self,
|
||||
uri: str,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta | None = None,
|
||||
) -> ResourceResult | mcp.types.CreateTaskResult:
|
||||
|
|
@ -1752,8 +1665,8 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
uri,
|
||||
resource_uri=uri,
|
||||
) as span:
|
||||
# Try concrete resources first (auth checked in get_resource)
|
||||
resource = await self.get_resource(uri, version=version)
|
||||
# Try concrete resources first (transforms + auth via _get_resource)
|
||||
resource = await self._get_resource(uri, version=version)
|
||||
if resource is not None:
|
||||
span.set_attributes(resource.get_span_attributes())
|
||||
if task_meta is not None and task_meta.fn_key is None:
|
||||
|
|
@ -1773,8 +1686,8 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
f"Error reading resource {uri!r}: {e}"
|
||||
) from e
|
||||
|
||||
# Try templates (auth checked in get_resource_template)
|
||||
template = await self.get_resource_template(uri, version=version)
|
||||
# Try templates (transforms + auth via _get_resource_template)
|
||||
template = await self._get_resource_template(uri, version=version)
|
||||
if template is None:
|
||||
if version is None:
|
||||
raise NotFoundError(f"Unknown resource: {uri!r}")
|
||||
|
|
@ -1803,7 +1716,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: None = None,
|
||||
) -> PromptResult: ...
|
||||
|
|
@ -1814,7 +1727,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta,
|
||||
) -> mcp.types.CreateTaskResult: ...
|
||||
|
|
@ -1824,7 +1737,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
name: str,
|
||||
arguments: dict[str, Any] | None = None,
|
||||
*,
|
||||
version: str | None = None,
|
||||
version: VersionSpec | None = None,
|
||||
run_middleware: bool = True,
|
||||
task_meta: TaskMeta | None = None,
|
||||
) -> PromptResult | mcp.types.CreateTaskResult:
|
||||
|
|
@ -1874,10 +1787,11 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
)
|
||||
|
||||
# Core logic: find and render prompt (providers queried in parallel)
|
||||
# Use _get_prompt to apply transforms (including visibility)
|
||||
with server_span(
|
||||
f"prompts/get {name}", "prompts/get", self.name, "prompt", name
|
||||
) as span:
|
||||
prompt = await self.get_prompt(name, version=version)
|
||||
prompt = await self._get_prompt(name, version=version)
|
||||
if prompt is None:
|
||||
raise NotFoundError(f"Unknown prompt: {name!r}")
|
||||
span.set_attributes(prompt.get_span_attributes())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -301,17 +301,19 @@ class TestToolLevelAuth:
|
|||
finally:
|
||||
auth_context_var.reset(tok)
|
||||
|
||||
async def test_get_tool_returns_none_without_auth(self):
|
||||
"""get_tool() checks auth and returns None for unauthorized tools."""
|
||||
async def test_get_tool_raises_without_auth(self):
|
||||
"""get_tool() checks auth and raises AuthorizationError for unauthorized tools."""
|
||||
from fastmcp.exceptions import AuthorizationError
|
||||
|
||||
mcp = FastMCP()
|
||||
|
||||
@mcp.tool(auth=require_auth)
|
||||
def protected_tool() -> str:
|
||||
return "protected"
|
||||
|
||||
# get_tool() returns None for unauthorized tools
|
||||
tool = await mcp.get_tool("protected_tool")
|
||||
assert tool is None
|
||||
# get_tool() raises AuthorizationError for unauthorized tools
|
||||
with pytest.raises(AuthorizationError, match="Unauthorized access to tool"):
|
||||
await mcp.get_tool("protected_tool")
|
||||
|
||||
async def test_get_tool_returns_tool_with_auth(self):
|
||||
mcp = FastMCP()
|
||||
|
|
|
|||
|
|
@ -351,10 +351,6 @@ class TestPromptEnabled:
|
|||
prompts = await mcp.get_prompts()
|
||||
assert len(prompts) == 0
|
||||
|
||||
# get_prompt() returns None for disabled prompts (Provider interface)
|
||||
prompt = await mcp.get_prompt("sample_prompt")
|
||||
assert prompt is None
|
||||
|
||||
async def test_prompt_toggle_enabled(self):
|
||||
mcp = FastMCP()
|
||||
|
||||
|
|
@ -381,8 +377,8 @@ class TestPromptEnabled:
|
|||
prompts = await mcp.get_prompts()
|
||||
assert len(prompts) == 0
|
||||
|
||||
# get_prompt() returns None for disabled prompts (Provider interface)
|
||||
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):
|
||||
|
|
@ -392,15 +388,15 @@ 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
|
||||
|
||||
# get_prompt() returns None for disabled prompts (Provider interface)
|
||||
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):
|
||||
|
|
@ -412,8 +408,8 @@ class TestPromptEnabled:
|
|||
|
||||
mcp.disable(keys=["prompt:sample_prompt@"])
|
||||
|
||||
# get_prompt() returns None for disabled prompts (Provider interface)
|
||||
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
|
||||
|
||||
|
||||
|
|
@ -458,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"
|
||||
|
||||
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"})
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
@ -483,12 +484,12 @@ class TestVersionFilter:
|
|||
assert tools[0].version == "3.0"
|
||||
|
||||
# Can request specific versions in range
|
||||
tool_v2 = await mcp.get_tool("add", version="2.0")
|
||||
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 - returns None
|
||||
assert await mcp.get_tool("add", version="1.0") is 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."""
|
||||
|
|
@ -520,13 +521,13 @@ class TestVersionFilter:
|
|||
assert tools[0].version == "2.5"
|
||||
|
||||
# Can request specific versions in range
|
||||
tool_v2 = await mcp.get_tool("calc", version="2.0")
|
||||
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 - return None
|
||||
assert await mcp.get_tool("calc", version="1.0") is None
|
||||
assert await mcp.get_tool("calc", version="3.0") is 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."""
|
||||
|
|
@ -576,7 +577,7 @@ 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."""
|
||||
"""get_tool() returns None if highest version is filtered out."""
|
||||
|
||||
from fastmcp.server.transforms import VersionFilter
|
||||
|
||||
|
|
@ -883,12 +884,12 @@ class TestMountedRangeFiltering:
|
|||
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")
|
||||
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 return None
|
||||
result = await parent.get_tool("child_calc", version="3.0")
|
||||
result = await parent.get_tool("child_calc", VersionSpec(eq="3.0"))
|
||||
assert result is None
|
||||
|
||||
|
||||
|
|
@ -930,7 +931,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
|
||||
|
||||
|
|
@ -1066,12 +1067,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
|
||||
|
||||
|
|
@ -1092,7 +1097,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):
|
||||
|
|
@ -1113,7 +1118,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"
|
||||
|
||||
|
|
@ -1130,7 +1135,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:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue