refactor: reverse visibility for get_tool/_get_tool methods

This commit is contained in:
Jeremiah Lowin 2026-01-17 14:38:07 -05:00
commit b13e1c7caf
14 changed files with 48 additions and 50 deletions

View file

@ -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 |

View file

@ -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"

View file

@ -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

View file

@ -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]

View file

@ -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:

View file

@ -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)

View file

@ -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."""

View file

@ -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.

View file

@ -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."""

View file

@ -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())

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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