|
8 | 8 | from collections.abc import Awaitable, Callable, Mapping, Sequence |
9 | 9 | from contextlib import AbstractAsyncContextManager, AsyncExitStack |
10 | 10 | from dataclasses import KW_ONLY, dataclass, field |
11 | | -from typing import TYPE_CHECKING, Any, Literal, TypeAlias, TypeVar, cast |
| 11 | +from typing import Any, Literal, TypeVar, cast |
12 | 12 |
|
13 | 13 | import anyio |
14 | 14 | import anyio.lowlevel |
|
40 | 40 | ServerCapabilities, |
41 | 41 | ) |
42 | 42 | from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, MODERN_PROTOCOL_VERSIONS |
43 | | -from typing_extensions import deprecated |
| 43 | +from typing_extensions import Protocol, deprecated |
44 | 44 |
|
45 | 45 | from mcp_client.client._input_required import DEFAULT_INPUT_REQUIRED_MAX_ROUNDS, run_input_required_driver |
46 | 46 | from mcp_client.client._probe import negotiate_auto |
|
61 | 61 | from mcp_client.client.streamable_http import streamable_http_client |
62 | 62 | from mcp_client.client.subscriptions import ServerEvent, Subscription |
63 | 63 | from mcp_client.client.subscriptions import listen as _listen |
64 | | -from mcp_client.shared.direct_dispatcher import create_direct_dispatcher_pair |
65 | 64 | from mcp_client.shared.dispatcher import Dispatcher, ProgressFnT |
66 | 65 | from mcp_client.shared.exceptions import MCPDeprecationWarning, MCPError |
67 | 66 | from mcp_client.shared.extension import validate_extension_identifier |
68 | 67 | from mcp_client.shared.jsonrpc_dispatcher import JSONRPCDispatcher |
69 | 68 | from mcp_client.shared.subscriptions import event_to_notification |
70 | 69 |
|
71 | | -if TYPE_CHECKING: |
72 | | - from mcp.server import Server |
73 | | - from mcp.server.mcpserver import MCPServer |
74 | | - |
75 | | - _InProcessServer: TypeAlias = Server[Any] | MCPServer |
76 | | -else: |
77 | | - # The full SDK binds its server types without making them a client dependency. |
78 | | - _InProcessServer = Any |
79 | | - |
80 | 70 | logger = logging.getLogger("mcp.client.client") |
81 | 71 |
|
82 | 72 | ConnectMode = Literal["legacy", "auto"] | str |
@@ -105,28 +95,10 @@ async def connect(exit_stack: AsyncExitStack, _mode: ConnectMode, _raise_excepti |
105 | 95 | return connect |
106 | 96 |
|
107 | 97 |
|
108 | | -def _connect_inproc(server: Server[Any]) -> _Connector: |
109 | | - """Connector for an in-process ``Server``: legacy mode drives the stream loop via |
110 | | - ``InMemoryTransport``; any other mode drives the modern per-request path through a |
111 | | - ``DirectDispatcher`` peer pair (no streams, no JSON-RPC framing, no initialize handshake).""" |
112 | | - |
113 | | - from mcp.server.runner import modern_on_request |
114 | | - from mcp_client.client._memory import InMemoryTransport |
115 | | - |
116 | | - async def connect(exit_stack: AsyncExitStack, mode: ConnectMode, raise_exceptions: bool) -> Dispatcher[Any]: |
117 | | - if mode == "legacy": |
118 | | - transport = InMemoryTransport(server, raise_exceptions=raise_exceptions) |
119 | | - read_stream, write_stream = await exit_stack.enter_async_context(transport) |
120 | | - return JSONRPCDispatcher(read_stream, write_stream) |
121 | | - lifespan_state = await exit_stack.enter_async_context(server.lifespan(server)) |
122 | | - client_disp, server_disp = create_direct_dispatcher_pair(raise_handler_exceptions=raise_exceptions) |
123 | | - tg = await exit_stack.enter_async_context(anyio.create_task_group()) |
124 | | - exit_stack.callback(server_disp.close) |
125 | | - on_request = modern_on_request(server, lifespan_state) |
126 | | - await tg.start(server_disp.run, on_request, _no_inbound_client_notifications) |
127 | | - return client_disp |
128 | | - |
129 | | - return connect |
| 98 | +class _InProcessServer(Protocol): |
| 99 | + async def __mcp_client_connect__( |
| 100 | + self, exit_stack: AsyncExitStack, mode: str, raise_exceptions: bool |
| 101 | + ) -> Dispatcher[Any]: ... |
130 | 102 |
|
131 | 103 |
|
132 | 104 | def _connected(value: _T | None) -> _T: |
@@ -189,17 +161,6 @@ def _synthesize_discover(protocol_version: str) -> types.DiscoverResult: |
189 | 161 | ) |
190 | 162 |
|
191 | 163 |
|
192 | | -async def _no_inbound_client_notifications(_dctx: Any, _method: str, _params: Mapping[str, Any] | None) -> None: |
193 | | - """Server-side inbound ``OnNotify`` for the modern in-process path — receives nothing. |
194 | | -
|
195 | | - At 2026-07-28 the spec defines no client→server notifications: ``initialized`` and |
196 | | - ``roots/list_changed`` are removed, and cancellation is structural (anyio scope cancel |
197 | | - through the direct await, not a notify). Server→client notifications (progress, log |
198 | | - messages) flow the other way via the per-request ``DispatchContext`` into the client's |
199 | | - callbacks, and are not seen here. |
200 | | - """ |
201 | | - |
202 | | - |
203 | 164 | @dataclass(frozen=True) |
204 | 165 | class _FoldedExtensions: |
205 | 166 | """`Client.extensions` instances folded into the shapes `ClientSession` consumes.""" |
@@ -278,7 +239,7 @@ class Client: |
278 | 239 | ```python |
279 | 240 | import asyncio |
280 | 241 |
|
281 | | - from mcp import Client |
| 242 | + from mcp_client import Client |
282 | 243 |
|
283 | 244 | async def main(): |
284 | 245 | async with Client("http://localhost:8000/mcp") as client: |
@@ -401,11 +362,7 @@ def __post_init__(self) -> None: |
401 | 362 | elif isinstance(srv, AbstractAsyncContextManager): |
402 | 363 | self._connect = _connect_transport(srv) |
403 | 364 | else: |
404 | | - from mcp.server.mcpserver import MCPServer |
405 | | - |
406 | | - if isinstance(srv, MCPServer): |
407 | | - srv = srv._lowlevel_server # pyright: ignore[reportPrivateUsage] |
408 | | - self._connect = _connect_inproc(cast("Server[Any]", srv)) |
| 365 | + self._connect = cast(_InProcessServer, srv).__mcp_client_connect__ |
409 | 366 |
|
410 | 367 | if self.cache is not None: |
411 | 368 | config = self.cache |
|
0 commit comments