|
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 Any, Literal, TypeVar, cast |
| 11 | +from typing import Any, Literal, TypeAlias, 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 Protocol, deprecated |
| 43 | +from typing_extensions import Protocol, deprecated, runtime_checkable |
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 |
@@ -95,12 +95,17 @@ async def connect(exit_stack: AsyncExitStack, _mode: ConnectMode, _raise_excepti |
95 | 95 | return connect |
96 | 96 |
|
97 | 97 |
|
98 | | -class _InProcessServer(Protocol): |
| 98 | +@runtime_checkable |
| 99 | +class _ServerConnector(Protocol): |
99 | 100 | async def __mcp_client_connect__( |
100 | 101 | self, exit_stack: AsyncExitStack, mode: str, raise_exceptions: bool |
101 | 102 | ) -> Dispatcher[Any]: ... |
102 | 103 |
|
103 | 104 |
|
| 105 | +# The full SDK rebinds this annotation alias, not the runtime-checkable protocol. |
| 106 | +_InProcessServer: TypeAlias = _ServerConnector |
| 107 | + |
| 108 | + |
104 | 109 | def _connected(value: _T | None) -> _T: |
105 | 110 | """Narrow a post-handshake session attribute from ``T | None`` to ``T``. |
106 | 111 |
|
@@ -355,14 +360,14 @@ def __post_init__(self) -> None: |
355 | 360 | self._folded_extensions = _fold_extensions(self.extensions) |
356 | 361 |
|
357 | 362 | srv = self.server |
358 | | - if isinstance(srv, str): |
| 363 | + if isinstance(srv, _ServerConnector): |
| 364 | + self._connect = srv.__mcp_client_connect__ |
| 365 | + elif isinstance(srv, str): |
359 | 366 | self._connect = _connect_transport(streamable_http_client(srv)) |
360 | 367 | elif isinstance(srv, StdioServerParameters): |
361 | 368 | self._connect = _connect_transport(stdio_client(srv)) |
362 | | - elif isinstance(srv, AbstractAsyncContextManager): |
363 | | - self._connect = _connect_transport(srv) |
364 | 369 | else: |
365 | | - self._connect = cast(_InProcessServer, srv).__mcp_client_connect__ |
| 370 | + self._connect = _connect_transport(srv) |
366 | 371 |
|
367 | 372 | if self.cache is not None: |
368 | 373 | config = self.cache |
|
0 commit comments