Merge pull request #1 from sandipan1/feature/load-server-config-client

This commit is contained in:
Sandipan Haldar 2025-04-26 02:15:16 +05:30 committed by GitHub
commit b9ff33b382
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 37 additions and 5 deletions

View file

@ -1 +0,0 @@

View file

@ -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,

View file

@ -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}")