mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 22:14:18 +02:00
refactor: reverse visibility for get_prompt/_get_prompt methods
This commit is contained in:
parent
24cf45bb64
commit
2c9cae9a98
8 changed files with 18 additions and 18 deletions
|
|
@ -295,7 +295,7 @@ class AuthMiddleware(Middleware):
|
|||
)
|
||||
|
||||
# Get prompt (component auth is checked in get_prompt, raises if unauthorized)
|
||||
prompt = await fastmcp.fastmcp._get_prompt(prompt_name)
|
||||
prompt = await fastmcp.fastmcp.get_prompt(prompt_name)
|
||||
if prompt is None:
|
||||
raise AuthorizationError(
|
||||
f"Authorization failed for prompt '{prompt_name}': prompt not found"
|
||||
|
|
|
|||
|
|
@ -217,7 +217,7 @@ class AggregateProvider(Provider):
|
|||
)
|
||||
return self._collect_list_results(results, "list_prompts")
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get prompt by name.
|
||||
|
|
@ -228,7 +228,7 @@ class AggregateProvider(Provider):
|
|||
If specified, returns highest version matching the spec from any provider.
|
||||
"""
|
||||
results = await gather(
|
||||
*[p._get_prompt(name, version) for p in self._providers],
|
||||
*[p.get_prompt(name, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
return self._get_highest_version_result(results, f"get_prompt({name!r})") # type: ignore[return-value]
|
||||
|
|
|
|||
|
|
@ -217,7 +217,7 @@ class Provider:
|
|||
|
||||
return await chain()
|
||||
|
||||
async def _get_prompt(
|
||||
async def get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get prompt by transformed name with all transforms applied.
|
||||
|
|
@ -228,7 +228,7 @@ class Provider:
|
|||
"""
|
||||
|
||||
async def base(n: str, version: VersionSpec | None = None) -> Prompt | None:
|
||||
return await self.get_prompt(n, version)
|
||||
return await self._get_prompt(n, version)
|
||||
|
||||
chain = base
|
||||
for transform in self.transforms:
|
||||
|
|
@ -342,12 +342,12 @@ class Provider:
|
|||
"""
|
||||
return []
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a specific prompt by name.
|
||||
|
||||
Default implementation filters list_prompts() and picks the highest version
|
||||
Default implementation filters _list_prompts() and picks the highest version
|
||||
matching the spec.
|
||||
|
||||
Args:
|
||||
|
|
|
|||
|
|
@ -578,16 +578,16 @@ class FastMCPProvider(Provider):
|
|||
raw_prompts = await self.server.get_prompts(run_middleware=True)
|
||||
return [FastMCPProviderPrompt.wrap(self.server, p) for p in raw_prompts]
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a prompt by name as a FastMCPProviderPrompt.
|
||||
|
||||
Passes the full VersionSpec to the nested server, which handles both
|
||||
exact version matching and range filtering. Uses _get_prompt to ensure
|
||||
exact version matching and range filtering. Uses get_prompt to ensure
|
||||
the nested server's transforms are applied.
|
||||
"""
|
||||
raw_prompt = await self.server._get_prompt(name, version)
|
||||
raw_prompt = await self.server.get_prompt(name, version)
|
||||
if raw_prompt is None:
|
||||
return None
|
||||
return FastMCPProviderPrompt.wrap(self.server, raw_prompt)
|
||||
|
|
|
|||
|
|
@ -215,12 +215,12 @@ class FileSystemProvider(LocalProvider):
|
|||
await self._ensure_loaded()
|
||||
return await super()._list_prompts()
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a prompt by name, reloading if in reload mode."""
|
||||
await self._ensure_loaded()
|
||||
return await super().get_prompt(name, version)
|
||||
return await super()._get_prompt(name, version)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"FileSystemProvider(root={self._root!r}, reload={self._reload})"
|
||||
|
|
|
|||
|
|
@ -591,7 +591,7 @@ class LocalProvider(Provider):
|
|||
if isinstance(v, Prompt) and self._is_component_enabled(v)
|
||||
]
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a prompt by name.
|
||||
|
|
|
|||
|
|
@ -1421,12 +1421,12 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
|
||||
return _dedupe_with_versions(authorized, lambda p: p.name)
|
||||
|
||||
async def get_prompt(
|
||||
async def _get_prompt(
|
||||
self, name: str, version: VersionSpec | None = None
|
||||
) -> Prompt | None:
|
||||
"""Get a prompt by name via aggregation from providers.
|
||||
|
||||
This is the raw lookup that Provider._get_prompt() wraps with transforms.
|
||||
This is the raw lookup that Provider.get_prompt() wraps with transforms.
|
||||
Aggregates from all sub-providers and applies component-level auth.
|
||||
|
||||
Args:
|
||||
|
|
@ -1438,7 +1438,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
"""
|
||||
# Aggregate from all sub-providers (each applies their own transforms)
|
||||
results = await gather(
|
||||
*[p._get_prompt(name, version) for p in self._providers],
|
||||
*[p.get_prompt(name, version) for p in self._providers],
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
|
|
@ -1791,7 +1791,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
|
|||
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())
|
||||
|
|
|
|||
|
|
@ -428,7 +428,7 @@ class TestProviderExecutionMethods:
|
|||
"""Test that default render_prompt uses get_prompt and renders it."""
|
||||
|
||||
class PromptProvider(Provider):
|
||||
async def list_prompts(self) -> Sequence[Prompt]:
|
||||
async def _list_prompts(self) -> Sequence[Prompt]:
|
||||
return [
|
||||
FunctionPrompt.from_function(
|
||||
fn=lambda name: f"Hello, {name}!",
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue