Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/fastmcp/server/providers/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ async def _get_tool(self, name: str) -> Tool | None:

from typing import TYPE_CHECKING

from fastmcp.server.providers.aggregate import AggregateProvider
from fastmcp.server.providers.base import Provider
from fastmcp.server.providers.fastmcp_provider import FastMCPProvider
from fastmcp.server.providers.filesystem import FileSystemProvider
Expand All @@ -37,6 +38,7 @@ async def _get_tool(self, name: str) -> Tool | None:
from fastmcp.server.providers.proxy import ProxyProvider as ProxyProvider

__all__ = [
"AggregateProvider",
"FastMCPProvider",
"FileSystemProvider",
"LocalProvider",
Expand Down
154 changes: 79 additions & 75 deletions src/fastmcp/server/providers/aggregate.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,9 @@
from fastmcp.server.providers import AggregateProvider

# Combine multiple providers into one
combined = AggregateProvider([provider1, provider2, provider3])
combined = AggregateProvider()
combined.add_provider(provider1)
combined.add_provider(provider2, namespace="api") # Tools become "api_foo"

# Use like any other provider
tools = await combined.list_tools()
Expand All @@ -21,18 +23,21 @@
import logging
from collections.abc import AsyncIterator, Sequence
from contextlib import AsyncExitStack, asynccontextmanager
from typing import TypeVar
from typing import TYPE_CHECKING, TypeVar

from fastmcp.exceptions import NotFoundError
from fastmcp.prompts.prompt import Prompt
from fastmcp.resources.resource import Resource
from fastmcp.resources.template import ResourceTemplate
from fastmcp.server.providers.base import Provider
from fastmcp.tools.tool import Tool
from fastmcp.server.transforms import Namespace
from fastmcp.utilities.async_utils import gather
from fastmcp.utilities.components import FastMCPComponent
from fastmcp.utilities.versions import VersionSpec, version_sort_key

if TYPE_CHECKING:
from fastmcp.prompts.prompt import Prompt
from fastmcp.resources.resource import Resource
from fastmcp.resources.template import ResourceTemplate
from fastmcp.tools.tool import Tool

logger = logging.getLogger(__name__)

T = TypeVar("T")
Expand All @@ -44,20 +49,59 @@ class AggregateProvider(Provider):
Components are aggregated from all providers. For get_* operations,
providers are queried in parallel and the highest version is returned.

When adding providers with a namespace, wrap_transform() is used to apply
the Namespace transform. This means namespace transformation is handled
by the wrapped provider, not by AggregateProvider.

Errors from individual providers are logged and skipped (graceful degradation).

This is useful when you want to combine custom providers without creating
a full FastMCP server.
Example:
```python
combined = AggregateProvider()
combined.add_provider(db_provider)
combined.add_provider(api_provider, namespace="api")
# db_provider's tools keep original names
# api_provider's tools become "api_foo", "api_bar", etc.
```
"""

def __init__(self, providers: Sequence[Provider]) -> None:
"""Initialize with a sequence of providers.
def __init__(self, providers: Sequence[Provider] | None = None) -> None:
"""Initialize with an optional sequence of providers.

Args:
providers: The providers to aggregate. Queried in order for lookups.
providers: Optional initial providers (without namespacing).
For namespaced providers, use add_provider() instead.
"""
super().__init__()
self._providers = list(providers)
self.providers: list[Provider] = list(providers or [])

def add_provider(self, provider: Provider, *, namespace: str = "") -> None:
"""Add a provider with optional namespace.

If the provider is a FastMCP server, it's automatically wrapped in
FastMCPProvider to ensure middleware is invoked correctly.

Args:
provider: The provider to add.
namespace: Optional namespace prefix. When set:
- Tools become "namespace_toolname"
- Resources become "protocol://namespace/path"
- Prompts become "namespace_promptname"
"""
# Import here to avoid circular imports
from fastmcp.server.server import FastMCP

# Auto-wrap FastMCP servers to ensure middleware is invoked
if isinstance(provider, FastMCP):
from fastmcp.server.providers.fastmcp_provider import FastMCPProvider

provider = FastMCPProvider(provider)

# Apply namespace via wrap_transform if specified
if namespace:
provider = provider.wrap_transform(Namespace(namespace))

self.providers.append(provider)

def _collect_list_results(
self, results: list[Sequence[T] | BaseException], operation: str
Expand All @@ -68,29 +112,12 @@ def _collect_list_results(
if isinstance(result, BaseException):
logger.debug(
f"Error during {operation} from provider "
f"{self._providers[i]}: {result}"
f"{self.providers[i]}: {result}"
)
continue
collected.extend(result)
return collected

def _get_first_result(
self, results: list[T | None | BaseException], operation: str
) -> T | None:
"""Get first successful non-None result, logging non-NotFoundError exceptions."""
for i, result in enumerate(results):
if isinstance(result, BaseException):
# NotFoundError is expected - don't log it
if not isinstance(result, NotFoundError):
logger.debug(
f"Error during {operation} from provider "
f"{self._providers[i]}: {result}"
)
continue
if result is not None:
return result
return None

def _get_highest_version_result(
self,
results: list[FastMCPComponent | None | BaseException],
Expand All @@ -107,7 +134,7 @@ def _get_highest_version_result(
if not isinstance(result, NotFoundError):
logger.debug(
f"Error during {operation} from provider "
f"{self._providers[i]}: {result}"
f"{self.providers[i]}: {result}"
)
continue
if result is not None:
Expand All @@ -117,32 +144,26 @@ def _get_highest_version_result(
return max(valid, key=version_sort_key)

def __repr__(self) -> str:
return f"AggregateProvider(providers={self._providers!r})"
return f"AggregateProvider(providers={self.providers!r})"

# -------------------------------------------------------------------------
# Tools
# -------------------------------------------------------------------------

async def _list_tools(self) -> Sequence[Tool]:
"""List all tools from all providers (with transforms applied)."""
"""List all tools from all providers."""
results = await gather(
*[p.list_tools() for p in self._providers],
*[p.list_tools() for p in self.providers],
return_exceptions=True,
)
return self._collect_list_results(results, "list_tools")

async def _get_tool(
self, name: str, version: VersionSpec | None = None
) -> Tool | None:
"""Get tool by name.

Args:
name: The tool name.
version: If None, returns highest version across all providers.
If specified, returns highest version matching the spec from any provider.
"""
"""Get tool by name from providers."""
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]
Expand All @@ -152,25 +173,19 @@ async def _get_tool(
# -------------------------------------------------------------------------

async def _list_resources(self) -> Sequence[Resource]:
"""List all resources from all providers (with transforms applied)."""
"""List all resources from all providers."""
results = await gather(
*[p.list_resources() for p in self._providers],
*[p.list_resources() for p in self.providers],
return_exceptions=True,
)
return self._collect_list_results(results, "list_resources")

async def _get_resource(
self, uri: str, version: VersionSpec | None = None
) -> Resource | None:
"""Get resource by URI.

Args:
uri: The resource URI.
version: If None, returns highest version across all providers.
If specified, returns highest version matching the spec from any provider.
"""
"""Get resource by URI from providers."""
results = await gather(
*[p.get_resource(uri, version) for p in self._providers],
*[p.get_resource(uri, version) for p in self.providers],
return_exceptions=True,
)
return self._get_highest_version_result(results, f"get_resource({uri!r})") # type: ignore[return-value]
Expand All @@ -180,25 +195,19 @@ async def _get_resource(
# -------------------------------------------------------------------------

async def _list_resource_templates(self) -> Sequence[ResourceTemplate]:
"""List all resource templates from all providers (with transforms applied)."""
"""List all resource templates from all providers."""
results = await gather(
*[p.list_resource_templates() for p in self._providers],
*[p.list_resource_templates() for p in self.providers],
return_exceptions=True,
)
return self._collect_list_results(results, "list_resource_templates")

async def _get_resource_template(
self, uri: str, version: VersionSpec | None = None
) -> ResourceTemplate | None:
"""Get resource template by URI.

Args:
uri: The template URI to match.
version: If None, returns highest version across all providers.
If specified, returns highest version matching the spec from any provider.
"""
"""Get resource template by URI from providers."""
results = await gather(
*[p.get_resource_template(uri, version) for p in self._providers],
*[p.get_resource_template(uri, version) for p in self.providers],
return_exceptions=True,
)
return self._get_highest_version_result(
Expand All @@ -210,25 +219,19 @@ async def _get_resource_template(
# -------------------------------------------------------------------------

async def _list_prompts(self) -> Sequence[Prompt]:
"""List all prompts from all providers (with transforms applied)."""
"""List all prompts from all providers."""
results = await gather(
*[p.list_prompts() for p in self._providers],
*[p.list_prompts() for p in self.providers],
return_exceptions=True,
)
return self._collect_list_results(results, "list_prompts")

async def _get_prompt(
self, name: str, version: VersionSpec | None = None
) -> Prompt | None:
"""Get prompt by name.

Args:
name: The prompt name.
version: If None, returns highest version across all providers.
If specified, returns highest version matching the spec from any provider.
"""
"""Get prompt by name from providers."""
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]
Expand All @@ -240,7 +243,8 @@ async def _get_prompt(
async def get_tasks(self) -> Sequence[FastMCPComponent]:
"""Get all task-eligible components from all providers."""
results = await gather(
*[p.get_tasks() for p in self._providers], return_exceptions=True
*[p.get_tasks() for p in self.providers],
return_exceptions=True,
)
return self._collect_list_results(results, "get_tasks")

Expand All @@ -252,6 +256,6 @@ async def get_tasks(self) -> Sequence[FastMCPComponent]:
async def lifespan(self) -> AsyncIterator[None]:
"""Combine lifespans of all providers."""
async with AsyncExitStack() as stack:
for provider in self._providers:
await stack.enter_async_context(provider.lifespan())
for p in self.providers:
await stack.enter_async_context(p.lifespan())
yield
32 changes: 32 additions & 0 deletions src/fastmcp/server/providers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,38 @@ def add_transform(self, transform: Transform) -> None:
"""
self._transforms.append(transform)

def wrap_transform(self, transform: Transform) -> Provider:
"""Return a new provider with this transform applied (immutable).

Unlike add_transform() which mutates this provider, wrap_transform()
returns a new provider that wraps this one. The original provider
is unchanged.

This is useful when you want to apply transforms without side effects,
such as adding the same provider to multiple aggregators with different
namespaces.

Args:
transform: The transform to apply.

Returns:
A new provider that wraps this one with the transform applied.

Example:
```python
from fastmcp.server.transforms import Namespace

provider = MyProvider()
namespaced = provider.wrap_transform(Namespace("api"))
# provider is unchanged
# namespaced returns tools as "api_toolname"
```
"""
# Import here to avoid circular imports
from fastmcp.server.providers.wrapped_provider import _WrappedProvider

return _WrappedProvider(self, transform)

# -------------------------------------------------------------------------
# Internal transform chain building
# -------------------------------------------------------------------------
Expand Down
Loading
Loading