From 648cbfa22244a52314535a6409bce3c0e42f4164 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Tue, 20 May 2025 14:38:10 -0400 Subject: [PATCH] Improved support for config dicts --- src/fastmcp/client/mcp_config.py | 74 ++++++++++++++++++++++++++++++++ src/fastmcp/client/transports.py | 54 +++++++++++------------ 2 files changed, 99 insertions(+), 29 deletions(-) create mode 100644 src/fastmcp/client/mcp_config.py diff --git a/src/fastmcp/client/mcp_config.py b/src/fastmcp/client/mcp_config.py new file mode 100644 index 000000000..d89eaf2f7 --- /dev/null +++ b/src/fastmcp/client/mcp_config.py @@ -0,0 +1,74 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Annotated, Any, Literal, TypeAlias + +from pydantic import Field +from pydantic.dataclasses import dataclass + +if TYPE_CHECKING: + from fastmcp.client.client import Client + from fastmcp.client.transports import ( + SSETransport, + StdioTransport, + StreamableHttpTransport, + ) + + +@dataclass +class LocalMCPServer: + command: str + args: list[str] + env: dict[str, Any] = Field(default_factory=dict) + cwd: str | None = None + + def to_transport(self) -> StdioTransport: + from fastmcp.client.transports import StdioTransport + + return StdioTransport( + command=self.command, + args=self.args, + env=self.env, + cwd=self.cwd, + ) + + +@dataclass +class RemoteMCPServer: + url: str + transport: Literal["http", "sse"] | 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) + else: + return SSETransport(self.url, headers=self.headers) + + +MCPServer: TypeAlias = LocalMCPServer | RemoteMCPServer + + +@dataclass +class MCPConfig: + mcp_servers: Annotated[dict[str, MCPServer], Field(alias="mcpServers")] + + @classmethod + def from_dict(cls, config: dict[str, Any]) -> MCPConfig: + return cls(mcp_servers=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() + } + + def to_clients(self) -> dict[str, Client]: + from fastmcp.client.client import Client + + return { + name: Client(transport=transport) + for name, transport in self.to_transports().items() + } diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index a5f0f95e0..b30a6e0d6 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 Any, TypedDict, cast +from typing import TYPE_CHECKING, Any, TypedDict, cast from urllib.parse import urlparse from mcp import ClientSession, StdioServerParameters @@ -24,9 +24,13 @@ 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 +if TYPE_CHECKING: + from fastmcp.client.mcp_config import MCPConfig + logger = get_logger(__name__) @@ -470,7 +474,13 @@ class FastMCPTransport(ClientTransport): def infer_transport( - transport: ClientTransport | FastMCPServer | AnyUrl | Path | dict[str, Any] | str, + transport: ClientTransport + | FastMCPServer + | AnyUrl + | Path + | MCPConfig + | dict[str, Any] + | str, ) -> ClientTransport: """ Infer the appropriate transport type from the given transport argument. @@ -481,6 +491,8 @@ 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 + # the transport is already a ClientTransport if isinstance(transport, ClientTransport): return transport @@ -511,34 +523,18 @@ def infer_transport( else: inferred_transport = StreamableHttpTransport(url=transport) - ## if the transport is a config dict - elif isinstance(transport, dict): - if "mcpServers" not in transport: - raise ValueError("Invalid transport dictionary: missing 'mcpServers' key") + # if the transport is a config dict or MCPConfig + elif isinstance(transport, dict | MCPConfig): + if isinstance(transport, dict): + config = MCPConfig.from_dict(transport) else: - server = transport["mcpServers"] - if len(list(server.keys())) > 1: - raise ValueError( - "Invalid transport dictionary: multiple servers found - only one expected" - ) - server_name = list(server.keys())[0] - # Stdio transport - if "command" in server[server_name] and "args" in server[server_name]: - inferred_transport = StdioTransport( - command=server[server_name]["command"], - args=server[server_name]["args"], - env=server[server_name].get("env", None), - cwd=server[server_name].get("cwd", None), - ) - - # HTTP transport - elif "url" in server: - inferred_transport = SSETransport( - url=server["url"], - headers=server.get("headers", None), - ) - - raise ValueError("Cannot determine transport type from dictionary") + config = transport + inferred_transports = config.to_transports() + if len(inferred_transports) > 1: + raise ValueError( + "Invalid transport dictionary: multiple servers found - only one expected" + ) + inferred_transport = list(inferred_transports.values())[0] # the transport is an unknown type else: