Merge pull request #623 from jlowin/infer-typing

Improve type inference from client transport
This commit is contained in:
Jeremiah Lowin 2025-05-28 21:17:37 -04:00 committed by GitHub
commit 7b879642ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 99 additions and 9 deletions

View file

@ -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,45 @@ 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
def __init__(
self,
transport: ClientTransport
transport: ClientTransportT
| FastMCP
| AnyUrl
| Path
@ -96,7 +143,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 = cast(ClientTransportT, infer_transport(transport))
self._session: ClientSession | None = None
self._exit_stack: AsyncExitStack | None = None
self._nesting_counter: int = 0

View file

@ -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"<MCPConfig(config='{self.config}')>"
@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

View file

@ -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.
"""