diff --git a/src/fastmcp/client/base.py b/src/fastmcp/client/base.py index 8b1378917..e69de29bb 100644 --- a/src/fastmcp/client/base.py +++ b/src/fastmcp/client/base.py @@ -1 +0,0 @@ - diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 7fdd21899..8509511d2 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -35,7 +35,7 @@ class Client: def __init__( self, - transport: ClientTransport | FastMCP | AnyUrl | Path | str, + transport: ClientTransport | FastMCP | AnyUrl | Path | dict[str, Any] | str, # Common args roots: RootsList | RootsHandler | None = None, sampling_handler: SamplingHandler | None = None, diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index 210286bf2..c8dc0d4cc 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -9,7 +9,7 @@ from pathlib import Path from typing import ( TypedDict, ) - +from typing import Any from exceptiongroup import BaseExceptionGroup, catch from mcp import ClientSession, McpError, StdioServerParameters from mcp.client.session import ( @@ -416,7 +416,7 @@ class FastMCPTransport(ClientTransport): def infer_transport( - transport: ClientTransport | FastMCPServer | AnyUrl | Path | str, + transport: ClientTransport | FastMCPServer | AnyUrl | Path | dict[str, Any] | str, ) -> ClientTransport: """ Infer the appropriate transport type from the given transport argument. @@ -449,7 +449,40 @@ def infer_transport( # the transport is a websocket URL elif isinstance(transport, AnyUrl | str) and str(transport).startswith("ws"): return WSTransport(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") + 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]: + return 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: + return SSETransport( + url=server["url"], + headers=server.get("headers", None), + ) + + # WebSocket transport + elif "ws_url" in server: + return WSTransport( + url=server["ws_url"], + ) + + raise ValueError("Cannot determine transport type from dictionary") + # the transport is an unknown type else: raise ValueError(f"Could not infer a valid transport from: {transport}")