From f5cddc321902976aeb8f126e6f67575baa38f2fe Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Tue, 20 May 2025 17:43:32 -0400 Subject: [PATCH] Formalize MCP Config --- src/fastmcp/client/transports.py | 36 +++++--- src/fastmcp/server/proxy.py | 8 ++ .../{client => utilities}/mcp_config.py | 58 +++++++----- tests/utilities/test_mcp_config.py | 92 +++++++++++++++++++ 4 files changed, 162 insertions(+), 32 deletions(-) rename src/fastmcp/{client => utilities}/mcp_config.py (54%) create mode 100644 tests/utilities/test_mcp_config.py diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index b30a6e0d6..dcd3450c8 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -6,7 +6,7 @@ import shutil import sys from collections.abc import AsyncIterator from pathlib import Path -from typing import TYPE_CHECKING, Any, TypedDict, cast +from typing import TYPE_CHECKING, Any, Literal, TypedDict, cast from urllib.parse import urlparse from mcp import ClientSession, StdioServerParameters @@ -24,12 +24,12 @@ from mcp.shared.memory import create_connected_server_and_client_session from pydantic import AnyUrl from typing_extensions import Unpack -from fastmcp.client.mcp_config import MCPConfig from fastmcp.server import FastMCP as FastMCPServer from fastmcp.utilities.logging import get_logger +from fastmcp.utilities.mcp_config import MCPConfig if TYPE_CHECKING: - from fastmcp.client.mcp_config import MCPConfig + from fastmcp.utilities.mcp_config import MCPConfig logger = get_logger(__name__) @@ -491,7 +491,7 @@ def infer_transport( For HTTP URLs, they are assumed to be Streamable HTTP URLs unless they end in `/sse`. """ - from fastmcp.client.mcp_config import MCPConfig + from fastmcp.utilities.mcp_config import MCPConfig # the transport is already a ClientTransport if isinstance(transport, ClientTransport): @@ -512,13 +512,8 @@ def infer_transport( # the transport is an http(s) URL elif isinstance(transport, AnyUrl | str) and str(transport).startswith("http"): - transport_str = str(transport) - # Parse out just the path portion to check for /sse - parsed_url = urlparse(transport_str) - path = parsed_url.path - - # Check if path contains /sse/ or ends with /sse - if "/sse/" in path or path.rstrip("/").endswith("/sse"): + inferred_transport_type = infer_transport_type_from_url(transport) + if inferred_transport_type == "sse": inferred_transport = SSETransport(url=transport) else: inferred_transport = StreamableHttpTransport(url=transport) @@ -542,3 +537,22 @@ def infer_transport( logger.debug(f"Inferred transport: {inferred_transport}") return inferred_transport + + +def infer_transport_type_from_url( + url: str | AnyUrl, +) -> Literal["streamable-http", "sse"]: + """ + Infer the appropriate transport type from the given URL. + """ + url = str(url) + if not url.startswith("http"): + raise ValueError(f"Invalid URL: {url}") + + parsed_url = urlparse(url) + path = parsed_url.path + + if "/sse/" in path or path.rstrip("/").endswith("/sse"): + return "sse" + else: + return "streamable-http" diff --git a/src/fastmcp/server/proxy.py b/src/fastmcp/server/proxy.py index 8f7123bab..9fcdcdb26 100644 --- a/src/fastmcp/server/proxy.py +++ b/src/fastmcp/server/proxy.py @@ -25,6 +25,7 @@ from fastmcp.server.context import Context from fastmcp.server.server import FastMCP from fastmcp.tools.tool import Tool from fastmcp.utilities.logging import get_logger +from fastmcp.utilities.mcp_config import MCPConfig if TYPE_CHECKING: from fastmcp.server import Context @@ -177,6 +178,13 @@ class FastMCPProxy(FastMCP): super().__init__(**kwargs) self.client = client + @classmethod + async def from_mcp_config(cls, config: MCPConfig | dict) -> FastMCPProxy: + if isinstance(config, dict): + config = MCPConfig.from_dict(config) + clients = config.to_clients() + return cls(client=clients[list(clients.keys())[0]]) + async def get_tools(self) -> dict[str, Tool]: tools = await super().get_tools() diff --git a/src/fastmcp/client/mcp_config.py b/src/fastmcp/utilities/mcp_config.py similarity index 54% rename from src/fastmcp/client/mcp_config.py rename to src/fastmcp/utilities/mcp_config.py index d89eaf2f7..7330bd65f 100644 --- a/src/fastmcp/client/mcp_config.py +++ b/src/fastmcp/utilities/mcp_config.py @@ -1,9 +1,9 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Annotated, Any, Literal, TypeAlias +from typing import TYPE_CHECKING, Any, Literal +from urllib.parse import urlparse -from pydantic import Field -from pydantic.dataclasses import dataclass +from pydantic import AnyUrl, BaseModel, Field if TYPE_CHECKING: from fastmcp.client.client import Client @@ -14,10 +14,28 @@ if TYPE_CHECKING: ) -@dataclass -class LocalMCPServer: +def infer_transport_type_from_url( + url: str | AnyUrl, +) -> Literal["streamable-http", "sse"]: + """ + Infer the appropriate transport type from the given URL. + """ + url = str(url) + if not url.startswith("http"): + raise ValueError(f"Invalid URL: {url}") + + parsed_url = urlparse(url) + path = parsed_url.path + + if "/sse/" in path or path.rstrip("/").endswith("/sse"): + return "sse" + else: + return "streamable-http" + + +class LocalMCPServer(BaseModel): command: str - args: list[str] + args: list[str] = Field(default_factory=list) env: dict[str, Any] = Field(default_factory=dict) cwd: str | None = None @@ -32,38 +50,36 @@ class LocalMCPServer: ) -@dataclass -class RemoteMCPServer: +class RemoteMCPServer(BaseModel): url: str - transport: Literal["http", "sse"] | None = None + transport: Literal["streamable-http", "sse", "http"] | None = None headers: dict[str, str] = Field(default_factory=dict) def to_transport(self) -> StreamableHttpTransport | SSETransport: from fastmcp.client.transports import SSETransport, StreamableHttpTransport - if self.transport in {"http", None}: - return StreamableHttpTransport(self.url, headers=self.headers) + if self.transport is None: + transport = infer_transport_type_from_url(self.url) else: + transport = self.transport + + if transport == "sse": return SSETransport(self.url, headers=self.headers) + else: + return StreamableHttpTransport(self.url, headers=self.headers) -MCPServer: TypeAlias = LocalMCPServer | RemoteMCPServer - - -@dataclass -class MCPConfig: - mcp_servers: Annotated[dict[str, MCPServer], Field(alias="mcpServers")] +class MCPConfig(BaseModel): + mcpServers: dict[str, LocalMCPServer | RemoteMCPServer] @classmethod def from_dict(cls, config: dict[str, Any]) -> MCPConfig: - return cls(mcp_servers=config.get("mcpServers", config)) + return cls(mcpServers=config.get("mcpServers", config)) def to_transports( self, ) -> dict[str, StdioTransport | StreamableHttpTransport | SSETransport]: - return { - name: server.to_transport() for name, server in self.mcp_servers.items() - } + return {name: server.to_transport() for name, server in self.mcpServers.items()} def to_clients(self) -> dict[str, Client]: from fastmcp.client.client import Client diff --git a/tests/utilities/test_mcp_config.py b/tests/utilities/test_mcp_config.py new file mode 100644 index 000000000..011a0574d --- /dev/null +++ b/tests/utilities/test_mcp_config.py @@ -0,0 +1,92 @@ +from fastmcp.client.transports import ( + SSETransport, + StdioTransport, + StreamableHttpTransport, +) +from fastmcp.utilities.mcp_config import LocalMCPServer, MCPConfig, RemoteMCPServer + + +def test_parse_single_stdio_config(): + config = { + "mcpServers": { + "test_server": { + "command": "echo", + "args": ["hello"], + } + } + } + mcp_config = MCPConfig.from_dict(config) + transport = mcp_config.mcpServers["test_server"].to_transport() + assert isinstance(transport, StdioTransport) + assert transport.command == "echo" + assert transport.args == ["hello"] + + +def test_parse_single_remote_config(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000", + } + } + } + mcp_config = MCPConfig.from_dict(config) + transport = mcp_config.mcpServers["test_server"].to_transport() + assert isinstance(transport, StreamableHttpTransport) + assert transport.url == "http://localhost:8000" + + +def test_parse_remote_config_with_transport(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000", + "transport": "sse", + } + } + } + mcp_config = MCPConfig.from_dict(config) + transport = mcp_config.mcpServers["test_server"].to_transport() + assert isinstance(transport, SSETransport) + assert transport.url == "http://localhost:8000" + + +def test_parse_remote_config_with_url_inference(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000/sse", + } + } + } + mcp_config = MCPConfig.from_dict(config) + transport = mcp_config.mcpServers["test_server"].to_transport() + assert isinstance(transport, SSETransport) + assert transport.url == "http://localhost:8000/sse" + + +def test_parse_multiple_servers(): + config = { + "mcpServers": { + "test_server": { + "url": "http://localhost:8000/sse", + }, + "test_server_2": { + "command": "echo", + "args": ["hello"], + "env": {"TEST": "test"}, + }, + } + } + mcp_config = MCPConfig.from_dict(config) + assert len(mcp_config.mcpServers) == 2 + assert isinstance(mcp_config.mcpServers["test_server"], RemoteMCPServer) + assert isinstance(mcp_config.mcpServers["test_server"].to_transport(), SSETransport) + + assert isinstance(mcp_config.mcpServers["test_server_2"], LocalMCPServer) + assert isinstance( + mcp_config.mcpServers["test_server_2"].to_transport(), StdioTransport + ) + assert mcp_config.mcpServers["test_server_2"].command == "echo" + assert mcp_config.mcpServers["test_server_2"].args == ["hello"] + assert mcp_config.mcpServers["test_server_2"].env == {"TEST": "test"}