mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-09 07:09:11 +02:00
refactor: reverse visibility for get_tool/_get_tool methods
This commit is contained in:
parent
5fe471dc60
commit
b13e1c7caf
14 changed files with 48 additions and 50 deletions
14
AGENTS.md
14
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 |
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue