From b13e1c7cafe2ae8f530d456f1e440e2c969f599d Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 14:38:07 -0500 Subject: [PATCH] refactor: reverse visibility for get_tool/_get_tool methods --- AGENTS.md | 14 ++++---- .../server/middleware/authorization.py | 2 +- src/fastmcp/server/providers/__init__.py | 2 +- src/fastmcp/server/providers/aggregate.py | 4 +-- src/fastmcp/server/providers/base.py | 8 ++--- .../server/providers/fastmcp_provider.py | 6 ++-- src/fastmcp/server/providers/filesystem.py | 4 +-- .../server/providers/local_provider.py | 2 +- .../server/providers/openapi/provider.py | 2 +- src/fastmcp/server/server.py | 10 +++--- tests/server/providers/test_local_provider.py | 4 +-- .../providers/test_transforming_provider.py | 6 ++-- tests/server/test_providers.py | 2 +- tests/server/test_versioning.py | 32 +++++++++---------- 14 files changed, 48 insertions(+), 50 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index 34a9e57db..1d6ff4e9b 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -6,23 +6,21 @@ FastMCP is a comprehensive Python framework (Python ≥3.10) for building Model ## Required Development Workflow -**CRITICAL**: Always run these commands in sequence before committing: +**CRITICAL**: Always run these commands in sequence before committing. ```bash uv sync # Install dependencies -uv run prek run --all-files # Ruff + Prettier + ty uv run pytest -n auto # Run full test suite ``` -**All three must pass** - this is enforced by CI. Alternative: `just build && just typecheck && just test` +In addition, you must pass static checks. This is generally done as a pre-commit hook with `prek` but you can run it manually with: + +```bash +uv run prek run --all-files # Ruff + Prettier + ty +``` **Tests must pass and lint/typing must be clean before committing.** -**Before creating a PR**, evaluate whether documentation needs updating: -- New features or APIs require corresponding docs -- Changed behavior should be reflected in existing docs -- Check `docs/` for affected pages - ## Repository Structure | Path | Purpose | diff --git a/src/fastmcp/server/middleware/authorization.py b/src/fastmcp/server/middleware/authorization.py index 7b377b478..1b696fe20 100644 --- a/src/fastmcp/server/middleware/authorization.py +++ b/src/fastmcp/server/middleware/authorization.py @@ -137,7 +137,7 @@ class AuthMiddleware(Middleware): ) # Get tool (component auth is checked in get_tool, raises if unauthorized) - tool = await fastmcp.fastmcp._get_tool(tool_name) + tool = await fastmcp.fastmcp.get_tool(tool_name) if tool is None: raise AuthorizationError( f"Authorization failed for tool '{tool_name}': tool not found" diff --git a/src/fastmcp/server/providers/__init__.py b/src/fastmcp/server/providers/__init__.py index 7001024a0..565266c1f 100644 --- a/src/fastmcp/server/providers/__init__.py +++ b/src/fastmcp/server/providers/__init__.py @@ -17,7 +17,7 @@ Example: rows = await self.db.fetch("SELECT * FROM tools") return [self._make_tool(row) for row in rows] - async def get_tool(self, name: str) -> Tool | None: + async def _get_tool(self, name: str) -> Tool | None: row = await self.db.fetchone("SELECT * FROM tools WHERE name = ?", name) return self._make_tool(row) if row else None diff --git a/src/fastmcp/server/providers/aggregate.py b/src/fastmcp/server/providers/aggregate.py index 861527ec4..ee22d9880 100644 --- a/src/fastmcp/server/providers/aggregate.py +++ b/src/fastmcp/server/providers/aggregate.py @@ -131,7 +131,7 @@ class AggregateProvider(Provider): ) return self._collect_list_results(results, "list_tools") - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get tool by name. @@ -142,7 +142,7 @@ class AggregateProvider(Provider): If specified, returns highest version matching the spec from any provider. """ results = await gather( - *[p._get_tool(name, version) for p in self._providers], + *[p.get_tool(name, version) for p in self._providers], return_exceptions=True, ) return self._get_highest_version_result(results, f"get_tool({name!r})") # type: ignore[return-value] diff --git a/src/fastmcp/server/providers/base.py b/src/fastmcp/server/providers/base.py index 7bb69db87..9a73d62ce 100644 --- a/src/fastmcp/server/providers/base.py +++ b/src/fastmcp/server/providers/base.py @@ -119,7 +119,7 @@ class Provider: return await chain() - async def _get_tool( + async def get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get tool by transformed name with all transforms applied. @@ -133,7 +133,7 @@ class Provider: """ async def base(n: str, version: VersionSpec | None = None) -> Tool | None: - return await self.get_tool(n, version) + return await self._get_tool(n, version) chain = base for transform in self.transforms: @@ -248,12 +248,12 @@ class Provider: """ return [] - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get a specific tool by name. - Default implementation filters list_tools() and picks the highest version + Default implementation filters _list_tools() and picks the highest version that matches the spec. Args: diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index 75c283db1..debdf70fc 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -492,16 +492,16 @@ class FastMCPProvider(Provider): raw_tools = await self.server.get_tools(run_middleware=True) return [FastMCPProviderTool.wrap(self.server, t) for t in raw_tools] - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """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. Uses _get_tool to ensure + exact version matching and range filtering. Uses get_tool to ensure the nested server's transforms are applied. """ - raw_tool = await self.server._get_tool(name, version) + raw_tool = await self.server.get_tool(name, version) if raw_tool is None: return None return FastMCPProviderTool.wrap(self.server, raw_tool) diff --git a/src/fastmcp/server/providers/filesystem.py b/src/fastmcp/server/providers/filesystem.py index fe4dba01f..d539a9ccc 100644 --- a/src/fastmcp/server/providers/filesystem.py +++ b/src/fastmcp/server/providers/filesystem.py @@ -179,12 +179,12 @@ class FileSystemProvider(LocalProvider): await self._ensure_loaded() return await super()._list_tools() - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get a tool by name, reloading if in reload mode.""" await self._ensure_loaded() - return await super().get_tool(name, version) + return await super()._get_tool(name, version) async def list_resources(self) -> Sequence[Resource]: """Return all resources, reloading if in reload mode.""" diff --git a/src/fastmcp/server/providers/local_provider.py b/src/fastmcp/server/providers/local_provider.py index 5584a91e1..17fafda68 100644 --- a/src/fastmcp/server/providers/local_provider.py +++ b/src/fastmcp/server/providers/local_provider.py @@ -500,7 +500,7 @@ class LocalProvider(Provider): if isinstance(v, Tool) and self._is_component_enabled(v) ] - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get a tool by name. diff --git a/src/fastmcp/server/providers/openapi/provider.py b/src/fastmcp/server/providers/openapi/provider.py index ca7fb3b7f..7330f70c9 100644 --- a/src/fastmcp/server/providers/openapi/provider.py +++ b/src/fastmcp/server/providers/openapi/provider.py @@ -353,7 +353,7 @@ class OpenAPIProvider(Provider): """Return all tools created from the OpenAPI spec.""" return list(self._tools.values()) - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get a tool by name.""" diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index f728595eb..c83266872 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -1125,12 +1125,12 @@ class FastMCP(Provider, Generic[LifespanResultT]): return _dedupe_with_versions(authorized, lambda t: t.name) - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: """Get a tool by name via aggregation from providers. - This is the raw lookup that Provider._get_tool() wraps with transforms. + This is the raw lookup that Provider.get_tool() wraps with transforms. Aggregates from all sub-providers and applies component-level auth. Args: @@ -1143,7 +1143,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Aggregate from all sub-providers (each applies their own transforms) results = await gather( - *[p._get_tool(name, version) for p in self._providers], + *[p.get_tool(name, version) for p in self._providers], return_exceptions=True, ) @@ -1177,7 +1177,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): return tool - # _get_tool is inherited from Provider - wraps get_tool() with transforms + # 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. @@ -1559,7 +1559,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): 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()) diff --git a/tests/server/providers/test_local_provider.py b/tests/server/providers/test_local_provider.py index e25a692b7..7dc69b3e3 100644 --- a/tests/server/providers/test_local_provider.py +++ b/tests/server/providers/test_local_provider.py @@ -578,7 +578,7 @@ class TestProviderToolTransformations: # Get tool through layer with call_next async def get_tool(name: str, version=None): - return await provider.get_tool(name, version) + return await provider._get_tool(name, version) tool = await layer.get_tool("transformed_tool", get_tool) assert tool is not None @@ -604,7 +604,7 @@ class TestProviderToolTransformations: ) async def get_tool(name: str, version=None): - return await provider.get_tool(name, version) + return await provider._get_tool(name, version) tool = await layer.get_tool("my_tool", get_tool) assert tool is not None diff --git a/tests/server/providers/test_transforming_provider.py b/tests/server/providers/test_transforming_provider.py index 8a7d6a254..28f89293b 100644 --- a/tests/server/providers/test_transforming_provider.py +++ b/tests/server/providers/test_transforming_provider.py @@ -159,7 +159,7 @@ class TestTransformReverseLookup: # Create call_next that delegates to provider async def get_tool(name: str, version=None): - return await provider.get_tool(name, version) + return await provider._get_tool(name, version) tool = await layer.get_tool("ns_my_tool", get_tool) @@ -178,7 +178,7 @@ class TestTransformReverseLookup: layer = ToolTransform({"original": ToolTransformConfig(name="renamed")}) async def get_tool(name: str, version=None): - return await provider.get_tool(name, version) + return await provider._get_tool(name, version) tool = await layer.get_tool("renamed", get_tool) @@ -216,7 +216,7 @@ class TestTransformReverseLookup: layer = Namespace("ns") async def get_tool(name: str, version=None): - return await provider.get_tool(name, version) + return await provider._get_tool(name, version) # Wrong namespace prefix assert await layer.get_tool("wrong_my_tool", get_tool) is None diff --git a/tests/server/test_providers.py b/tests/server/test_providers.py index 275053775..2bdaacbff 100644 --- a/tests/server/test_providers.py +++ b/tests/server/test_providers.py @@ -52,7 +52,7 @@ class SimpleToolProvider(Provider): self.list_tools_call_count += 1 return self._tools - async def get_tool( + async def _get_tool( self, name: str, version: VersionSpec | None = None ) -> Tool | None: self.get_tool_call_count += 1 diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index f2b0713dc..c78164aa9 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -483,13 +483,13 @@ class TestVersionFilter: assert len(tools) == 1 assert tools[0].version == "3.0" - # Can request specific versions in range (use _get_tool to apply transforms) - tool_v2 = await mcp._get_tool("add", VersionSpec(eq="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 - returns None - assert await mcp._get_tool("add", VersionSpec(eq="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,14 +520,14 @@ class TestVersionFilter: assert len(tools) == 1 assert tools[0].version == "2.5" - # Can request specific versions in range (use _get_tool to apply transforms) - tool_v2 = await mcp._get_tool("calc", VersionSpec(eq="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 - 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 + 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.""" @@ -589,8 +589,8 @@ class TestVersionFilter: mcp.add_transform(VersionFilter(version_lt="3.0")) - # Tool exists but is filtered out - returns None (use _get_tool to apply transforms) - assert await mcp._get_tool("only_v5") is None + # 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.""" @@ -831,8 +831,8 @@ class TestMountedVersionFiltering: tools = await parent.get_tools() assert len(tools) == 0 - # _get_tool should also return None (respects filter, applies transforms) - assert await parent._get_tool("child_high_version_tool") is None + # get_tool should also return None (respects filter, applies transforms) + assert await parent.get_tool("child_high_version_tool") is None class TestMountedRangeFiltering: @@ -857,8 +857,8 @@ class TestMountedRangeFiltering: parent.add_transform(VersionFilter(version_lt="2.0")) # Should return v1.0 (the highest version that matches <2.0) - # Use _get_tool to apply transforms - 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" @@ -884,13 +884,13 @@ class TestMountedRangeFiltering: parent.mount(child, "child") parent.add_transform(VersionFilter(version_gte="1.0", version_lt="3.0")) - # Request specific version within range (use _get_tool to apply transforms) - tool = await parent._get_tool("child_calc", VersionSpec(eq="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 return None - result = await parent._get_tool("child_calc", VersionSpec(eq="3.0")) + result = await parent.get_tool("child_calc", VersionSpec(eq="3.0")) assert result is None