mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 05:54:19 +02:00
Merge pull request #623 from jlowin/infer-typing
Improve type inference from client transport
This commit is contained in:
commit
7b879642ee
3 changed files with 99 additions and 9 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue