Jeremiah Lowin commited on
Commit
fa6e614
·
1 Parent(s): 582a7ae

Improve type inference from client transport

Browse files
src/fastmcp/client/client.py CHANGED
@@ -1,7 +1,7 @@
1
  import datetime
2
  from contextlib import AsyncExitStack, asynccontextmanager
3
  from pathlib import Path
4
- from typing import Any, cast
5
 
6
  import anyio
7
  import mcp.types
@@ -28,7 +28,18 @@ from fastmcp.server import FastMCP
28
  from fastmcp.utilities.exceptions import get_catch_handlers
29
  from fastmcp.utilities.mcp_config import MCPConfig
30
 
31
- from .transports import ClientTransport, SessionKwargs, infer_transport
 
 
 
 
 
 
 
 
 
 
 
32
 
33
  __all__ = [
34
  "Client",
@@ -41,7 +52,7 @@ __all__ = [
41
  ]
42
 
43
 
44
- class Client:
45
  """
46
  MCP client that delegates connection management to a Transport instance.
47
 
@@ -78,9 +89,47 @@ class Client:
78
  ```
79
  """
80
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
81
  def __init__(
82
  self,
83
- transport: ClientTransport
84
  | FastMCP
85
  | AnyUrl
86
  | Path
@@ -96,7 +145,7 @@ class Client:
96
  timeout: datetime.timedelta | float | int | None = None,
97
  init_timeout: datetime.timedelta | float | int | None = None,
98
  ):
99
- self.transport = infer_transport(transport)
100
  self._session: ClientSession | None = None
101
  self._exit_stack: AsyncExitStack | None = None
102
  self._nesting_counter: int = 0
 
1
  import datetime
2
  from contextlib import AsyncExitStack, asynccontextmanager
3
  from pathlib import Path
4
+ from typing import Any, Generic, cast, overload
5
 
6
  import anyio
7
  import mcp.types
 
28
  from fastmcp.utilities.exceptions import get_catch_handlers
29
  from fastmcp.utilities.mcp_config import MCPConfig
30
 
31
+ from .transports import (
32
+ ClientTransportT,
33
+ FastMCP1Server,
34
+ FastMCPTransport,
35
+ MCPConfigTransport,
36
+ NodeStdioTransport,
37
+ PythonStdioTransport,
38
+ SessionKwargs,
39
+ SSETransport,
40
+ StreamableHttpTransport,
41
+ infer_transport,
42
+ )
43
 
44
  __all__ = [
45
  "Client",
 
52
  ]
53
 
54
 
55
+ class Client(Generic[ClientTransportT]):
56
  """
57
  MCP client that delegates connection management to a Transport instance.
58
 
 
89
  ```
90
  """
91
 
92
+ @overload
93
+ def __new__(
94
+ cls,
95
+ transport: ClientTransportT,
96
+ **kwargs: Any,
97
+ ) -> "Client[ClientTransportT]": ...
98
+
99
+ @overload
100
+ def __new__(
101
+ cls, transport: AnyUrl, **kwargs
102
+ ) -> "Client[SSETransport|StreamableHttpTransport]": ...
103
+
104
+ @overload
105
+ def __new__(
106
+ cls, transport: FastMCP | FastMCP1Server, **kwargs
107
+ ) -> "Client[FastMCPTransport]": ...
108
+
109
+ @overload
110
+ def __new__(
111
+ cls, transport: Path, **kwargs
112
+ ) -> "Client[PythonStdioTransport|NodeStdioTransport]": ...
113
+
114
+ @overload
115
+ def __new__(
116
+ cls, transport: MCPConfig | dict[str, Any], **kwargs
117
+ ) -> "Client[MCPConfigTransport]": ...
118
+
119
+ @overload
120
+ def __new__(
121
+ cls, transport: str, **kwargs
122
+ ) -> "Client[PythonStdioTransport|NodeStdioTransport|SSETransport|StreamableHttpTransport]": ...
123
+
124
+ def __new__(cls, transport, **kwargs) -> "Client":
125
+ instance = super().__new__(cls)
126
+ return instance
127
+
128
+ transport: ClientTransportT
129
+
130
  def __init__(
131
  self,
132
+ transport: ClientTransportT
133
  | FastMCP
134
  | AnyUrl
135
  | Path
 
145
  timeout: datetime.timedelta | float | int | None = None,
146
  init_timeout: datetime.timedelta | float | int | None = None,
147
  ):
148
+ self.transport = infer_transport(transport) # type: ignore
149
  self._session: ClientSession | None = None
150
  self._exit_stack: AsyncExitStack | None = None
151
  self._nesting_counter: int = 0
src/fastmcp/client/transports.py CHANGED
@@ -6,7 +6,7 @@ import shutil
6
  import sys
7
  from collections.abc import AsyncIterator
8
  from pathlib import Path
9
- from typing import TYPE_CHECKING, Any, TypedDict, cast
10
 
11
  from mcp import ClientSession, StdioServerParameters
12
  from mcp.client.session import (
@@ -35,6 +35,9 @@ if TYPE_CHECKING:
35
 
36
  logger = get_logger(__name__)
37
 
 
 
 
38
 
39
  class SessionKwargs(TypedDict, total=False):
40
  """Keyword arguments for the MCP ClientSession constructor."""
@@ -575,6 +578,44 @@ class MCPConfigTransport(ClientTransport):
575
  return f"<MCPConfig(config='{self.config}')>"
576
 
577
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
578
  def infer_transport(
579
  transport: ClientTransport
580
  | FastMCPServer
 
6
  import sys
7
  from collections.abc import AsyncIterator
8
  from pathlib import Path
9
+ from typing import TYPE_CHECKING, Any, TypedDict, TypeVar, cast, overload
10
 
11
  from mcp import ClientSession, StdioServerParameters
12
  from mcp.client.session import (
 
35
 
36
  logger = get_logger(__name__)
37
 
38
+ # TypeVar for preserving specific ClientTransport subclass types
39
+ ClientTransportT = TypeVar("ClientTransportT", bound="ClientTransport")
40
+
41
 
42
  class SessionKwargs(TypedDict, total=False):
43
  """Keyword arguments for the MCP ClientSession constructor."""
 
578
  return f"<MCPConfig(config='{self.config}')>"
579
 
580
 
581
+ @overload
582
+ def infer_transport(transport: ClientTransportT) -> ClientTransportT: ...
583
+
584
+
585
+ @overload
586
+ def infer_transport(transport: FastMCPServer) -> FastMCPTransport: ...
587
+
588
+
589
+ @overload
590
+ def infer_transport(transport: FastMCP1Server) -> FastMCPTransport: ...
591
+
592
+
593
+ @overload
594
+ def infer_transport(transport: MCPConfig) -> MCPConfigTransport: ...
595
+
596
+
597
+ @overload
598
+ def infer_transport(transport: dict[str, Any]) -> MCPConfigTransport: ...
599
+
600
+
601
+ @overload
602
+ def infer_transport(
603
+ transport: AnyUrl,
604
+ ) -> SSETransport | StreamableHttpTransport: ...
605
+
606
+
607
+ @overload
608
+ def infer_transport(
609
+ transport: str,
610
+ ) -> (
611
+ PythonStdioTransport | NodeStdioTransport | SSETransport | StreamableHttpTransport
612
+ ): ...
613
+
614
+
615
+ @overload
616
+ def infer_transport(transport: Path) -> PythonStdioTransport | NodeStdioTransport: ...
617
+
618
+
619
  def infer_transport(
620
  transport: ClientTransport
621
  | FastMCPServer
src/fastmcp/server/server.py CHANGED
@@ -62,7 +62,7 @@ from fastmcp.utilities.mcp_config import MCPConfig
62
 
63
  if TYPE_CHECKING:
64
  from fastmcp.client import Client
65
- from fastmcp.client.transports import ClientTransport
66
  from fastmcp.server.openapi import ComponentFn as OpenAPIComponentFn
67
  from fastmcp.server.openapi import FastMCPOpenAPI, RouteMap
68
  from fastmcp.server.openapi import RouteMapFn as OpenAPIRouteMapFn
@@ -1288,7 +1288,7 @@ class FastMCP(Generic[LifespanResultT]):
1288
  @classmethod
1289
  def as_proxy(
1290
  cls,
1291
- backend: Client
1292
  | ClientTransport
1293
  | FastMCP[Any]
1294
  | AnyUrl
@@ -1316,7 +1316,9 @@ class FastMCP(Generic[LifespanResultT]):
1316
  return FastMCPProxy(client=client, **settings)
1317
 
1318
  @classmethod
1319
- def from_client(cls, client: Client, **settings: Any) -> FastMCPProxy:
 
 
1320
  """
1321
  Create a FastMCP proxy server from a FastMCP client.
1322
  """
 
62
 
63
  if TYPE_CHECKING:
64
  from fastmcp.client import Client
65
+ from fastmcp.client.transports import ClientTransport, ClientTransportT
66
  from fastmcp.server.openapi import ComponentFn as OpenAPIComponentFn
67
  from fastmcp.server.openapi import FastMCPOpenAPI, RouteMap
68
  from fastmcp.server.openapi import RouteMapFn as OpenAPIRouteMapFn
 
1288
  @classmethod
1289
  def as_proxy(
1290
  cls,
1291
+ backend: Client[ClientTransportT]
1292
  | ClientTransport
1293
  | FastMCP[Any]
1294
  | AnyUrl
 
1316
  return FastMCPProxy(client=client, **settings)
1317
 
1318
  @classmethod
1319
+ def from_client(
1320
+ cls, client: Client[ClientTransportT], **settings: Any
1321
+ ) -> FastMCPProxy:
1322
  """
1323
  Create a FastMCP proxy server from a FastMCP client.
1324
  """