diff --git a/fastmcp_remote/fastmcp_remote/cli.py b/fastmcp_remote/fastmcp_remote/cli.py index 3e3f766b1..22786abd3 100644 --- a/fastmcp_remote/fastmcp_remote/cli.py +++ b/fastmcp_remote/fastmcp_remote/cli.py @@ -234,7 +234,12 @@ def build_transport(config: RemoteConfig) -> SSETransport | StreamableHttpTransp async def run(config: RemoteConfig) -> None: client = Client(build_transport(config)) - server = create_proxy(client, name="fastmcp-remote") + server = create_proxy( + client, + name="fastmcp-remote", + provider_error_strategy="raise", + validate_on_initialize=True, + ) if config.ignore_tools: server.add_transform(IgnoreTools(config.ignore_tools)) await server.run_async( diff --git a/fastmcp_slim/fastmcp/server/providers/aggregate.py b/fastmcp_slim/fastmcp/server/providers/aggregate.py index 5d6b8ce01..766d64846 100644 --- a/fastmcp_slim/fastmcp/server/providers/aggregate.py +++ b/fastmcp_slim/fastmcp/server/providers/aggregate.py @@ -23,7 +23,7 @@ from __future__ import annotations import logging from collections.abc import AsyncIterator, Sequence from contextlib import AsyncExitStack, asynccontextmanager -from typing import TYPE_CHECKING, TypeVar +from typing import TYPE_CHECKING, Literal, TypeVar from fastmcp.exceptions import NotFoundError from fastmcp.server.providers.base import Provider @@ -41,6 +41,7 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) T = TypeVar("T") +ProviderErrorStrategy = Literal["warn", "raise"] class AggregateProvider(Provider): @@ -53,7 +54,9 @@ class AggregateProvider(Provider): 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). + Errors from individual providers are logged and skipped by default. Set + ``provider_error_strategy="raise"`` to fail the aggregate operation when + any provider fails. Example: ```python @@ -65,14 +68,23 @@ class AggregateProvider(Provider): ``` """ - def __init__(self, providers: Sequence[Provider] | None = None) -> None: + def __init__( + self, + providers: Sequence[Provider] | None = None, + *, + provider_error_strategy: ProviderErrorStrategy = "warn", + ) -> None: """Initialize with an optional sequence of providers. Args: providers: Optional initial providers (without namespacing). For namespaced providers, use add_provider() instead. + provider_error_strategy: How provider errors should affect aggregate + operations. ``"warn"`` logs and skips failed providers. + ``"raise"`` propagates the first provider error. """ super().__init__() + self.provider_error_strategy = provider_error_strategy self.providers: list[Provider] = list(providers or []) def add_provider(self, provider: Provider, *, namespace: str = "") -> None: @@ -121,6 +133,8 @@ class AggregateProvider(Provider): seen_keys: dict[str, int] = {} for i, result in enumerate(results): if isinstance(result, BaseException): + if self.provider_error_strategy == "raise": + raise result logger.warning( f"Error during {operation} from provider " f"{self.providers[i]}: {result}" @@ -153,6 +167,8 @@ class AggregateProvider(Provider): for i, result in enumerate(results): if isinstance(result, BaseException): if not isinstance(result, NotFoundError): + if self.provider_error_strategy == "raise": + raise result logger.warning( f"Error during {operation} from provider " f"{self.providers[i]}: {result}" @@ -197,6 +213,8 @@ class AggregateProvider(Provider): ) for r in results: if isinstance(r, BaseException): + if self.provider_error_strategy == "raise": + raise r continue if r is not None: return r @@ -210,6 +228,8 @@ class AggregateProvider(Provider): ) for r in results: if isinstance(r, BaseException): + if self.provider_error_strategy == "raise": + raise r continue if r is not None: return r diff --git a/fastmcp_slim/fastmcp/server/providers/proxy.py b/fastmcp_slim/fastmcp/server/providers/proxy.py index b7c84d5da..08f84615d 100644 --- a/fastmcp_slim/fastmcp/server/providers/proxy.py +++ b/fastmcp_slim/fastmcp/server/providers/proxy.py @@ -14,6 +14,8 @@ from collections.abc import Awaitable, Callable, Sequence from typing import TYPE_CHECKING, Any, cast from urllib.parse import quote +import anyio +import httpx import mcp.types from mcp import ServerSession from mcp.client.session import ClientSession @@ -42,6 +44,8 @@ from fastmcp.resources import Resource, ResourceTemplate from fastmcp.resources.base import ResourceContent, ResourceResult from fastmcp.server.context import Context from fastmcp.server.dependencies import get_context +from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext +from fastmcp.server.providers.aggregate import ProviderErrorStrategy from fastmcp.server.providers.base import Provider from fastmcp.server.server import FastMCP from fastmcp.server.tasks.config import TaskConfig @@ -61,6 +65,46 @@ logger = get_logger(__name__) ClientFactoryT = Callable[[], Client] | Callable[[], Awaitable[Client]] +def _proxy_upstream_error(error: Exception) -> McpError: + return McpError( + mcp.types.ErrorData( + code=mcp.types.INTERNAL_ERROR, + message=str(error), + ) + ) + + +class ProxyInitializeMiddleware(Middleware): + def __init__(self, proxy: FastMCPProxy) -> None: + self.proxy = proxy + + async def on_initialize( + self, + context: MiddlewareContext[mcp.types.InitializeRequest], + call_next: CallNext[ + mcp.types.InitializeRequest, + mcp.types.InitializeResult | None, + ], + ) -> mcp.types.InitializeResult | None: + client = await self.proxy._get_client() + try: + async with client: + pass + except McpError: + raise + except ( + RuntimeError, + TimeoutError, + httpx.HTTPError, + anyio.ClosedResourceError, + anyio.EndOfStream, + anyio.BrokenResourceError, + ) as error: + raise _proxy_upstream_error(error) from error + + return await call_next(context) + + # ----------------------------------------------------------------------------- # Proxy Component Classes # ----------------------------------------------------------------------------- @@ -836,6 +880,8 @@ class FastMCPProxy(FastMCP): self, *, client_factory: ClientFactoryT, + provider_error_strategy: ProviderErrorStrategy = "warn", + validate_on_initialize: bool = False, **kwargs, ): """Initialize the proxy server. @@ -847,12 +893,38 @@ class FastMCPProxy(FastMCP): client_factory: A callable that returns a Client instance when called. This gives you full control over session creation and reuse. Can be either a synchronous or asynchronous function. + provider_error_strategy: How provider errors should affect aggregate + operations. Defaults to ``"warn"`` for compatibility; use + ``"raise"`` when the proxy should surface upstream failures. + validate_on_initialize: If true, connect to the upstream server during + the incoming MCP initialize request. **kwargs: Additional settings for the FastMCP server. """ super().__init__(**kwargs) + self.provider_error_strategy = provider_error_strategy self.client_factory = client_factory provider: Provider = ProxyProvider(client_factory) self.add_provider(provider) + if validate_on_initialize: + self.middleware.append(ProxyInitializeMiddleware(self)) + self._setup_proxy_ping_handler() + + async def _get_client(self) -> Client: + client = self.client_factory() + if inspect.isawaitable(client): + client = cast(Client, await client) + return client + + def _setup_proxy_ping_handler(self) -> None: + async def ping_remote( + _request: mcp.types.PingRequest, + ) -> mcp.types.ServerResult: + client = await self._get_client() + async with client: + await client.ping() + return mcp.types.ServerResult(mcp.types.EmptyResult()) + + self._mcp_server.request_handlers[mcp.types.PingRequest] = ping_remote # ----------------------------------------------------------------------------- diff --git a/tests/server/providers/proxy/test_proxy_server.py b/tests/server/providers/proxy/test_proxy_server.py index 36557a96f..8e90cdd9e 100644 --- a/tests/server/providers/proxy/test_proxy_server.py +++ b/tests/server/providers/proxy/test_proxy_server.py @@ -27,6 +27,8 @@ from fastmcp.tools.base import ToolResult from fastmcp.tools.tool_transform import ( ToolTransformConfig, ) +from fastmcp.utilities.http import find_available_port +from fastmcp.utilities.tests import run_server_async USERS = [ {"id": "1", "name": "Alice", "active": True}, @@ -245,6 +247,58 @@ async def test_proxy_with_async_client_factory(): assert client.transport.url == "http://example.com/mcp/" +async def test_proxy_ping_forwards_to_remote_server(fastmcp_server): + proxy = create_proxy(fastmcp_server) + + async with Client(proxy) as client: + assert await client.ping() is True + + +async def test_proxy_ping_surfaces_wrong_remote_path(): + remote = FastMCP("remote") + async with run_server_async(remote, transport="http") as url: + proxy = create_proxy(StreamableHttpTransport(url.removesuffix("/mcp"))) + + async with Client(proxy) as client: + with pytest.raises(McpError, match="Session terminated"): + await client.ping() + + +async def test_proxy_initialize_surfaces_remote_connection_error(): + port = find_available_port() + proxy = create_proxy( + StreamableHttpTransport(f"http://127.0.0.1:{port}/mcp"), + validate_on_initialize=True, + ) + + with pytest.raises(McpError, match="Client failed to connect"): + async with Client(proxy): + pass + + +async def test_proxy_list_tools_surfaces_remote_connection_error(): + port = find_available_port() + proxy = create_proxy( + StreamableHttpTransport(f"http://127.0.0.1:{port}/mcp"), + provider_error_strategy="raise", + ) + + with pytest.raises(RuntimeError, match="Client failed to connect"): + await proxy.list_tools() + + +async def test_proxy_list_tools_client_surfaces_remote_connection_error(): + port = find_available_port() + proxy = create_proxy( + StreamableHttpTransport(f"http://127.0.0.1:{port}/mcp"), + provider_error_strategy="raise", + ) + + with pytest.raises(McpError, match="Client failed to connect"): + async with Client(proxy) as client: + await client.list_tools() + + class TestTools: async def test_get_tools(self, proxy_server): tools = await proxy_server.list_tools() diff --git a/tests/server/providers/test_base_provider.py b/tests/server/providers/test_base_provider.py index 0e8898a3f..38db55e08 100644 --- a/tests/server/providers/test_base_provider.py +++ b/tests/server/providers/test_base_provider.py @@ -2,6 +2,9 @@ from typing import Any +import pytest + +from fastmcp.server.providers.aggregate import AggregateProvider from fastmcp.server.providers.base import Provider from fastmcp.server.tasks.config import TaskConfig from fastmcp.server.transforms import Namespace @@ -29,6 +32,11 @@ class SimpleProvider(Provider): return self._tools +class FailingProvider(Provider): + async def _list_tools(self) -> list[Tool]: + raise RuntimeError("provider unavailable") + + class TestBaseProviderGetTasks: """Tests for Provider.get_tasks() base implementation.""" @@ -89,3 +97,19 @@ class TestBaseProviderGetTasks: assert len(tasks) == 1 assert tasks[0].name == "api_my_tool" + + +class TestAggregateProviderErrors: + async def test_provider_errors_warn_by_default(self): + aggregate = AggregateProvider([FailingProvider()]) + + assert await aggregate.list_tools() == [] + + async def test_provider_errors_can_raise(self): + aggregate = AggregateProvider( + [FailingProvider()], + provider_error_strategy="raise", + ) + + with pytest.raises(RuntimeError, match="provider unavailable"): + await aggregate.list_tools()