diff --git a/src/fastmcp/server/middleware/authorization.py b/src/fastmcp/server/middleware/authorization.py index 451c509b5..7b377b478 100644 --- a/src/fastmcp/server/middleware/authorization.py +++ b/src/fastmcp/server/middleware/authorization.py @@ -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" diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index e2d55804a..ed4f12ec4 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -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 diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index b3b61b8a2..df8f051cb 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -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()) diff --git a/src/fastmcp/server/tasks/requests.py b/src/fastmcp/server/tasks/requests.py index ff23f2f6c..61286d831 100644 --- a/src/fastmcp/server/tasks/requests.py +++ b/src/fastmcp/server/tasks/requests.py @@ -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 diff --git a/tests/server/auth/test_authorization.py b/tests/server/auth/test_authorization.py index 79d4e5e3d..6a64dff3f 100644 --- a/tests/server/auth/test_authorization.py +++ b/tests/server/auth/test_authorization.py @@ -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() diff --git a/tests/server/providers/test_local_provider_prompts.py b/tests/server/providers/test_local_provider_prompts.py index 0445401bb..666ecb62d 100644 --- a/tests/server/providers/test_local_provider_prompts.py +++ b/tests/server/providers/test_local_provider_prompts.py @@ -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" diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index a6e451c81..d6c61f78d 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -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: