From 432dd4dd2d32c9c66158d9516480c565e78d4d36 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 28 May 2025 21:05:34 -0400 Subject: [PATCH 1/2] Improve type inference from client transport --- src/fastmcp/client/client.py | 59 +++++++++++++++++++++++++++++--- src/fastmcp/client/transports.py | 43 ++++++++++++++++++++++- src/fastmcp/server/server.py | 8 +++-- 3 files changed, 101 insertions(+), 9 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 853f68c79..dbb5b1793 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -1,7 +1,7 @@ import datetime from contextlib import AsyncExitStack, asynccontextmanager from pathlib import Path -from typing import Any, cast +from typing import Any, Generic, cast, overload import anyio import mcp.types @@ -28,7 +28,18 @@ from fastmcp.server import FastMCP from fastmcp.utilities.exceptions import get_catch_handlers from fastmcp.utilities.mcp_config import MCPConfig -from .transports import ClientTransport, SessionKwargs, infer_transport +from .transports import ( + ClientTransportT, + FastMCP1Server, + FastMCPTransport, + MCPConfigTransport, + NodeStdioTransport, + PythonStdioTransport, + SessionKwargs, + SSETransport, + StreamableHttpTransport, + infer_transport, +) __all__ = [ "Client", @@ -41,7 +52,7 @@ __all__ = [ ] -class Client: +class Client(Generic[ClientTransportT]): """ MCP client that delegates connection management to a Transport instance. @@ -78,9 +89,47 @@ class Client: ``` """ + @overload + def __new__( + cls, + transport: ClientTransportT, + **kwargs: Any, + ) -> "Client[ClientTransportT]": ... + + @overload + def __new__( + cls, transport: AnyUrl, **kwargs + ) -> "Client[SSETransport|StreamableHttpTransport]": ... + + @overload + def __new__( + cls, transport: FastMCP | FastMCP1Server, **kwargs + ) -> "Client[FastMCPTransport]": ... + + @overload + def __new__( + cls, transport: Path, **kwargs + ) -> "Client[PythonStdioTransport|NodeStdioTransport]": ... + + @overload + def __new__( + cls, transport: MCPConfig | dict[str, Any], **kwargs + ) -> "Client[MCPConfigTransport]": ... + + @overload + def __new__( + cls, transport: str, **kwargs + ) -> "Client[PythonStdioTransport|NodeStdioTransport|SSETransport|StreamableHttpTransport]": ... + + def __new__(cls, transport, **kwargs) -> "Client": + instance = super().__new__(cls) + return instance + + transport: ClientTransportT + def __init__( self, - transport: ClientTransport + transport: ClientTransportT | FastMCP | AnyUrl | Path @@ -96,7 +145,7 @@ class Client: timeout: datetime.timedelta | float | int | None = None, init_timeout: datetime.timedelta | float | int | None = None, ): - self.transport = infer_transport(transport) + self.transport = infer_transport(transport) # type: ignore self._session: ClientSession | None = None self._exit_stack: AsyncExitStack | None = None self._nesting_counter: int = 0 diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index e43ab5760..9a7268340 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, TypedDict, TypeVar, cast, overload from mcp import ClientSession, StdioServerParameters from mcp.client.session import ( @@ -35,6 +35,9 @@ if TYPE_CHECKING: logger = get_logger(__name__) +# TypeVar for preserving specific ClientTransport subclass types +ClientTransportT = TypeVar("ClientTransportT", bound="ClientTransport") + class SessionKwargs(TypedDict, total=False): """Keyword arguments for the MCP ClientSession constructor.""" @@ -575,6 +578,44 @@ class MCPConfigTransport(ClientTransport): return f"" +@overload +def infer_transport(transport: ClientTransportT) -> ClientTransportT: ... + + +@overload +def infer_transport(transport: FastMCPServer) -> FastMCPTransport: ... + + +@overload +def infer_transport(transport: FastMCP1Server) -> FastMCPTransport: ... + + +@overload +def infer_transport(transport: MCPConfig) -> MCPConfigTransport: ... + + +@overload +def infer_transport(transport: dict[str, Any]) -> MCPConfigTransport: ... + + +@overload +def infer_transport( + transport: AnyUrl, +) -> SSETransport | StreamableHttpTransport: ... + + +@overload +def infer_transport( + transport: str, +) -> ( + PythonStdioTransport | NodeStdioTransport | SSETransport | StreamableHttpTransport +): ... + + +@overload +def infer_transport(transport: Path) -> PythonStdioTransport | NodeStdioTransport: ... + + def infer_transport( transport: ClientTransport | FastMCPServer diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index f5e35916b..183c4927f 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -62,7 +62,7 @@ from fastmcp.utilities.mcp_config import MCPConfig if TYPE_CHECKING: from fastmcp.client import Client - from fastmcp.client.transports import ClientTransport + from fastmcp.client.transports import ClientTransport, ClientTransportT from fastmcp.server.openapi import ComponentFn as OpenAPIComponentFn from fastmcp.server.openapi import FastMCPOpenAPI, RouteMap from fastmcp.server.openapi import RouteMapFn as OpenAPIRouteMapFn @@ -1288,7 +1288,7 @@ class FastMCP(Generic[LifespanResultT]): @classmethod def as_proxy( cls, - backend: Client + backend: Client[ClientTransportT] | ClientTransport | FastMCP[Any] | AnyUrl @@ -1316,7 +1316,9 @@ class FastMCP(Generic[LifespanResultT]): return FastMCPProxy(client=client, **settings) @classmethod - def from_client(cls, client: Client, **settings: Any) -> FastMCPProxy: + def from_client( + cls, client: Client[ClientTransportT], **settings: Any + ) -> FastMCPProxy: """ Create a FastMCP proxy server from a FastMCP client. """ From a3450b4b2585d7ffceda2ed37361cea9088226c6 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 28 May 2025 21:15:47 -0400 Subject: [PATCH 2/2] Improve type inference --- src/fastmcp/client/client.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index dbb5b1793..a33969c8a 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -125,8 +125,6 @@ class Client(Generic[ClientTransportT]): instance = super().__new__(cls) return instance - transport: ClientTransportT - def __init__( self, transport: ClientTransportT @@ -145,7 +143,7 @@ class Client(Generic[ClientTransportT]): timeout: datetime.timedelta | float | int | None = None, init_timeout: datetime.timedelta | float | int | None = None, ): - self.transport = infer_transport(transport) # type: ignore + self.transport = cast(ClientTransportT, infer_transport(transport)) self._session: ClientSession | None = None self._exit_stack: AsyncExitStack | None = None self._nesting_counter: int = 0