955 lines
42 KiB
Python
955 lines
42 KiB
Python
"""Unified MCP Client that wraps ClientSession with transport management."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import logging
|
|
import uuid
|
|
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
|
from contextlib import AbstractAsyncContextManager, AsyncExitStack
|
|
from dataclasses import KW_ONLY, dataclass, field
|
|
from typing import Any, Literal, TypeVar, cast
|
|
|
|
import anyio
|
|
import anyio.lowlevel
|
|
import mcp_types as types
|
|
from mcp_types import (
|
|
INVALID_PARAMS,
|
|
CacheableResult,
|
|
CallToolResult,
|
|
CompleteResult,
|
|
EmptyResult,
|
|
ErrorData,
|
|
GetPromptResult,
|
|
Implementation,
|
|
InputRequest,
|
|
InputRequiredResult,
|
|
InputResponse,
|
|
InputResponses,
|
|
ListPromptsResult,
|
|
ListResourcesResult,
|
|
ListResourceTemplatesResult,
|
|
ListToolsResult,
|
|
LoggingLevel,
|
|
PaginatedRequestParams,
|
|
PromptReference,
|
|
ReadResourceResult,
|
|
RequestParamsMeta,
|
|
ResourceTemplateReference,
|
|
Result,
|
|
ServerCapabilities,
|
|
)
|
|
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, MODERN_PROTOCOL_VERSIONS
|
|
from typing_extensions import deprecated
|
|
|
|
from mcp.client._input_required import DEFAULT_INPUT_REQUIRED_MAX_ROUNDS, run_input_required_driver
|
|
from mcp.client._memory import InMemoryTransport
|
|
from mcp.client._probe import negotiate_auto
|
|
from mcp.client._transport import Transport
|
|
from mcp.client.caching import CacheConfig, CacheMode, ClientResponseCache, InMemoryResponseCacheStore
|
|
from mcp.client.extension import ClaimContext, ClientExtension, NotificationBinding, ResultClaim
|
|
from mcp.client.session import (
|
|
ClientRequestContext,
|
|
ClientSession,
|
|
ElicitationFnT,
|
|
IncomingMessage,
|
|
ListRootsFnT,
|
|
LoggingFnT,
|
|
MessageHandlerFnT,
|
|
SamplingFnT,
|
|
)
|
|
from mcp.client.stdio import StdioServerParameters, stdio_client
|
|
from mcp.client.streamable_http import streamable_http_client
|
|
from mcp.client.subscriptions import ServerEvent, Subscription
|
|
from mcp.client.subscriptions import listen as _listen
|
|
from mcp.server import Server
|
|
from mcp.server.mcpserver import MCPServer
|
|
from mcp.server.runner import modern_on_request
|
|
from mcp.shared.direct_dispatcher import create_direct_dispatcher_pair
|
|
from mcp.shared.dispatcher import Dispatcher, ProgressFnT
|
|
from mcp.shared.exceptions import MCPDeprecationWarning, MCPError
|
|
from mcp.shared.extension import validate_extension_identifier
|
|
from mcp.shared.jsonrpc_dispatcher import JSONRPCDispatcher
|
|
from mcp.shared.subscriptions import event_to_notification
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
ConnectMode = Literal["legacy", "auto"] | str
|
|
"""``mode=`` value: ``"legacy"`` (initialize handshake), ``"auto"`` (discover, fall back to
|
|
initialize), or a modern protocol-version string (adopt directly). The ``str`` arm is for
|
|
forward-compat; ``Client.__post_init__`` rejects anything outside that set at construction."""
|
|
|
|
_T = TypeVar("_T")
|
|
_ResultT = TypeVar("_ResultT")
|
|
_CacheableT = TypeVar("_CacheableT", bound=CacheableResult)
|
|
|
|
_Connector = Callable[[AsyncExitStack, ConnectMode, bool], Awaitable["Dispatcher[Any]"]]
|
|
"""Resolved at ``__post_init__`` from the shape of ``server`` alone: enter whatever resources
|
|
are needed onto the exit stack and hand back the ``Dispatcher`` ``ClientSession`` will drive.
|
|
``mode`` and ``raise_exceptions`` are passed at call time so they're read at the same moment
|
|
``__aenter__`` reads them for the handshake step."""
|
|
|
|
|
|
def _connect_transport(transport: Transport) -> _Connector:
|
|
"""Connector for the stream-backed paths (URL, user-supplied ``Transport``)."""
|
|
|
|
async def connect(exit_stack: AsyncExitStack, _mode: ConnectMode, _raise_exceptions: bool) -> Dispatcher[Any]:
|
|
read_stream, write_stream = await exit_stack.enter_async_context(transport)
|
|
return JSONRPCDispatcher(read_stream, write_stream)
|
|
|
|
return connect
|
|
|
|
|
|
def _connect_inproc(server: Server[Any]) -> _Connector:
|
|
"""Connector for an in-process ``Server``: legacy mode drives the stream loop via
|
|
``InMemoryTransport``; any other mode drives the modern per-request path through a
|
|
``DirectDispatcher`` peer pair (no streams, no JSON-RPC framing, no initialize handshake)."""
|
|
|
|
async def connect(exit_stack: AsyncExitStack, mode: ConnectMode, raise_exceptions: bool) -> Dispatcher[Any]:
|
|
if mode == "legacy":
|
|
transport = InMemoryTransport(server, raise_exceptions=raise_exceptions)
|
|
read_stream, write_stream = await exit_stack.enter_async_context(transport)
|
|
return JSONRPCDispatcher(read_stream, write_stream)
|
|
lifespan_state = await exit_stack.enter_async_context(server.lifespan(server))
|
|
client_disp, server_disp = create_direct_dispatcher_pair(raise_handler_exceptions=raise_exceptions)
|
|
tg = await exit_stack.enter_async_context(anyio.create_task_group())
|
|
exit_stack.callback(server_disp.close)
|
|
on_request = modern_on_request(server, lifespan_state)
|
|
await tg.start(server_disp.run, on_request, _no_inbound_client_notifications)
|
|
return client_disp
|
|
|
|
return connect
|
|
|
|
|
|
def _connected(value: _T | None) -> _T:
|
|
"""Narrow a post-handshake session attribute from ``T | None`` to ``T``.
|
|
|
|
``Client.__aenter__`` only assigns ``_session`` after the handshake succeeds, so inside
|
|
``async with Client(...)`` these attributes are always populated; the ``.session`` gate
|
|
raises before this is reached otherwise. The guard exists for pyright, not runtime.
|
|
"""
|
|
if value is None: # pragma: no cover
|
|
raise RuntimeError("Client must be used within an async context manager")
|
|
return value
|
|
|
|
|
|
def _strip_userinfo(url: str) -> str:
|
|
"""Drop any userinfo from the URL's authority component; byte-exact otherwise.
|
|
|
|
Credentials must not enter cache-key material; any further normalization could merge distinct servers.
|
|
"""
|
|
# Pure text, no urlsplit: it strips embedded tab/CR/LF before parsing, which would misalign slices.
|
|
sep = url.find("//")
|
|
if sep == -1:
|
|
return url
|
|
start = sep + 2
|
|
end = len(url)
|
|
for delimiter in "/?#":
|
|
if (found := url.find(delimiter, start)) != -1:
|
|
end = min(end, found)
|
|
authority = url[start:end]
|
|
if "@" not in authority:
|
|
return url
|
|
return url[:start] + authority.rpartition("@")[2] + url[end:]
|
|
|
|
|
|
def _evicting_message_handler(cache: ClientResponseCache, user_handler: MessageHandlerFnT | None) -> MessageHandlerFnT:
|
|
"""Wrap the session message handler with cache eviction on server notifications."""
|
|
|
|
async def handler(message: IncomingMessage) -> None:
|
|
if isinstance(message, types.ServerNotification):
|
|
try:
|
|
await cache.evict_for_notification(message)
|
|
except Exception: # boundary: eviction reaches user store code; a cache fault must not block delivery
|
|
logger.exception("Response cache eviction failed; the notification is still delivered")
|
|
if user_handler is not None:
|
|
await user_handler(message)
|
|
else:
|
|
# Mirrors ClientSession's default handler (session._default_message_handler).
|
|
await anyio.lowlevel.checkpoint()
|
|
|
|
return handler
|
|
|
|
|
|
def _synthesize_discover(protocol_version: str) -> types.DiscoverResult:
|
|
return types.DiscoverResult(
|
|
supported_versions=[protocol_version],
|
|
capabilities=types.ServerCapabilities(),
|
|
result_type="complete",
|
|
ttl_ms=0,
|
|
cache_scope="public",
|
|
)
|
|
|
|
|
|
async def _no_inbound_client_notifications(_dctx: Any, _method: str, _params: Mapping[str, Any] | None) -> None:
|
|
"""Server-side inbound ``OnNotify`` for the modern in-process path — receives nothing.
|
|
|
|
At 2026-07-28 the spec defines no client→server notifications: ``initialized`` and
|
|
``roots/list_changed`` are removed, and cancellation is structural (anyio scope cancel
|
|
through the direct await, not a notify). Server→client notifications (progress, log
|
|
messages) flow the other way via the per-request ``DispatchContext`` into the client's
|
|
callbacks, and are not seen here.
|
|
"""
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _FoldedExtensions:
|
|
"""`Client.extensions` instances folded into the shapes `ClientSession` consumes."""
|
|
|
|
ad: dict[str, dict[str, Any]] | None
|
|
claims: dict[str, tuple[ResultClaim[Any], ...]] | None
|
|
bindings: tuple[NotificationBinding[Any], ...] | None
|
|
by_model: Mapping[type[Result], ResultClaim[Any]]
|
|
|
|
|
|
def _fold_extensions(extensions: Sequence[ClientExtension] | None) -> _FoldedExtensions:
|
|
"""Fold extension contributions at construction, naming both owners on duplicate tags or methods."""
|
|
if isinstance(extensions, Mapping):
|
|
raise TypeError(
|
|
"extensions= takes a sequence of ClientExtension instances. The mapping form was "
|
|
"replaced: use advertise(identifier, settings) for advertise-only entries"
|
|
)
|
|
if not extensions:
|
|
return _FoldedExtensions(ad=None, claims=None, bindings=None, by_model={})
|
|
ad: dict[str, dict[str, Any]] = {}
|
|
claims: dict[str, tuple[ResultClaim[Any], ...]] = {}
|
|
bindings: list[NotificationBinding[Any]] = []
|
|
by_model: dict[type[Result], ResultClaim[Any]] = {}
|
|
claim_owners: dict[str, str] = {}
|
|
binding_owners: dict[str, str] = {}
|
|
for extension in extensions:
|
|
identifier = getattr(extension, "identifier", None)
|
|
if identifier is None:
|
|
raise ValueError(
|
|
f"{type(extension).__name__} has no `identifier`; a ClientExtension must set the "
|
|
"`identifier` class attribute (or assign one in `__init__`) before it can be used"
|
|
)
|
|
validate_extension_identifier(identifier, owner=type(extension).__name__)
|
|
if identifier in ad:
|
|
raise ValueError(f"extension identifier {identifier!r} is passed more than once")
|
|
ad[identifier] = extension.settings()
|
|
extension_claims = tuple(extension.claims())
|
|
for claim in extension_claims:
|
|
tag = claim.result_type
|
|
if tag in claim_owners:
|
|
owner = claim_owners[tag]
|
|
both = (
|
|
f"extension {identifier!r} claims"
|
|
if owner == identifier
|
|
else (f"extensions {owner!r} and {identifier!r} both claim")
|
|
)
|
|
raise ValueError(f"{both} resultType {tag!r}; a wire tag can have only one resolver")
|
|
claim_owners[tag] = identifier
|
|
# Each model pins its result_type Literal to one tag, so this index cannot collide.
|
|
by_model[claim.model] = claim
|
|
if extension_claims:
|
|
claims[identifier] = extension_claims
|
|
for binding in extension.notifications():
|
|
if binding.method in binding_owners:
|
|
owner = binding_owners[binding.method]
|
|
both = (
|
|
f"extension {identifier!r} binds"
|
|
if owner == identifier
|
|
else (f"extensions {owner!r} and {identifier!r} both bind")
|
|
)
|
|
raise ValueError(f"{both} notification method {binding.method!r}; a method can have only one observer")
|
|
binding_owners[binding.method] = identifier
|
|
bindings.append(binding)
|
|
return _FoldedExtensions(ad=ad, claims=claims or None, bindings=tuple(bindings) or None, by_model=by_model)
|
|
|
|
|
|
@dataclass
|
|
class Client:
|
|
"""A high-level MCP client for connecting to MCP servers.
|
|
|
|
Pass a URL string (Streamable HTTP), a `StdioServerParameters` (launch the command as a
|
|
subprocess and talk over its stdin/stdout), any `Transport`, or - in tests - a `Server` or
|
|
`MCPServer` instance to connect to it in-process.
|
|
|
|
Example:
|
|
```python
|
|
from mcp.client import Client
|
|
from mcp.server.mcpserver import MCPServer
|
|
|
|
server = MCPServer("test")
|
|
|
|
@server.tool()
|
|
def add(a: int, b: int) -> int:
|
|
return a + b
|
|
|
|
async def main():
|
|
async with Client(server) as client:
|
|
result = await client.call_tool("add", {"a": 1, "b": 2})
|
|
|
|
asyncio.run(main())
|
|
```
|
|
"""
|
|
|
|
server: Server[Any] | MCPServer | Transport | StdioServerParameters | str
|
|
"""The MCP server to connect to.
|
|
|
|
If the server is a URL string, it will be used as the URL for a `streamable_http_client` transport.
|
|
If the server is a `StdioServerParameters`, the command is launched with `stdio_client`.
|
|
If the server is a `Transport` instance, it will be used directly.
|
|
If the server is a `Server` or `MCPServer` instance, it will be connected in-process.
|
|
"""
|
|
|
|
_: KW_ONLY
|
|
|
|
# TODO(Marcelo): When do `raise_exceptions=True` actually raises?
|
|
raise_exceptions: bool = False
|
|
"""Whether to raise exceptions from the server."""
|
|
|
|
read_timeout_seconds: float | None = None
|
|
"""Timeout for read operations."""
|
|
|
|
sampling_callback: SamplingFnT | None = None
|
|
"""Callback for handling sampling requests."""
|
|
|
|
sampling_capabilities: types.SamplingCapability | None = None
|
|
"""Sampling sub-capabilities (e.g. tools) declared alongside `sampling_callback`; no effect without it."""
|
|
|
|
list_roots_callback: ListRootsFnT | None = None
|
|
"""Callback for handling list roots requests."""
|
|
|
|
logging_callback: LoggingFnT | None = None
|
|
"""Callback for handling logging notifications."""
|
|
|
|
log_level: LoggingLevel | None = None
|
|
"""The log level to opt in to on 2026-07-28+ connections (deprecated logging feature, SEP-2577).
|
|
|
|
Modern (2026-07-28+) servers send `notifications/message` only for requests that opt in by
|
|
carrying `io.modelcontextprotocol/logLevel` in `_meta`, and only at or above that level. Setting
|
|
this stamps that opt-in on every request; `None` (the default) means no opt-in, so no log
|
|
messages arrive - a `logging_callback` alone is not an opt-in. No effect on handshake-era
|
|
connections, where the deprecated `logging/setLevel` request governs delivery instead. A
|
|
per-request `_meta` entry with the same key overrides this default."""
|
|
|
|
# TODO(Marcelo): Why do we have both "callback" and "handler"?
|
|
message_handler: MessageHandlerFnT | None = None
|
|
"""Callback for handling raw messages."""
|
|
|
|
client_info: Implementation | None = None
|
|
"""Client implementation info to send to server."""
|
|
|
|
mode: ConnectMode = "auto"
|
|
"""How to negotiate the protocol version.
|
|
|
|
'auto' (the default) probes `server/discover` and falls back to the initialize handshake on legacy servers;
|
|
for an in-process `Server`/`MCPServer` it dispatches directly without JSON-RPC framing. 'legacy' forces the
|
|
initialize handshake (byte-identical pre-2026 behavior). A modern protocol-version string (e.g. '2026-07-28')
|
|
adopts that version directly without a probe — supply `prior_discover` to reuse a known DiscoverResult, or
|
|
omit it to synthesize a minimal one."""
|
|
|
|
prior_discover: types.DiscoverResult | None = None
|
|
"""A previously-obtained DiscoverResult to install via .adopt() when mode is a version pin.
|
|
Ignored when mode='legacy'."""
|
|
|
|
elicitation_callback: ElicitationFnT | None = None
|
|
"""Callback for handling elicitation requests."""
|
|
|
|
input_required_max_rounds: int = DEFAULT_INPUT_REQUIRED_MAX_ROUNDS
|
|
"""Cap on `InputRequiredResult` retry rounds before `call_tool` / `get_prompt` /
|
|
`read_resource` give up. Use `client.session.<method>(..., allow_input_required=True)`
|
|
to drive the loop manually instead."""
|
|
|
|
extensions: Sequence[ClientExtension] | None = None
|
|
"""Opt-in client extensions (SEP-2133).
|
|
|
|
Each instance contributes its capability ad, its result claims (resolved
|
|
transparently by `call_tool`), and its notification bindings. For an
|
|
ad-only entry use `mcp.client.advertise(identifier, settings)`."""
|
|
|
|
cache: CacheConfig | None = field(default_factory=CacheConfig)
|
|
"""Client-side response caching for the SEP-2549 cacheable methods (2026-07-28).
|
|
|
|
The default `CacheConfig()` honors server `ttlMs`/`cacheScope` hints with a
|
|
per-client in-memory store; pass a customized `CacheConfig`, or `None` to
|
|
disable. The cacheable verbs take a per-call `cache_mode` (see `CacheMode`);
|
|
calls carrying `meta` always reach the server. A `CacheConfig` with a custom
|
|
`store` requires `target_id` when the server is not a URL (no identity can be
|
|
derived)."""
|
|
|
|
_entered: bool = field(init=False, default=False)
|
|
_session: ClientSession | None = field(init=False, default=None)
|
|
_exit_stack: AsyncExitStack | None = field(init=False, default=None)
|
|
_connect: _Connector = field(init=False, repr=False, compare=False)
|
|
_response_cache: ClientResponseCache | None = field(init=False, default=None, repr=False, compare=False)
|
|
_folded_extensions: _FoldedExtensions = field(init=False, repr=False, compare=False)
|
|
|
|
def __post_init__(self) -> None:
|
|
if self.mode not in ("legacy", "auto") and self.mode not in MODERN_PROTOCOL_VERSIONS:
|
|
hint = (
|
|
f" ({self.mode!r} is a handshake-era version; use mode='legacy')"
|
|
if self.mode in HANDSHAKE_PROTOCOL_VERSIONS
|
|
else ""
|
|
)
|
|
raise ValueError(
|
|
f"mode must be 'legacy', 'auto', or one of {list(MODERN_PROTOCOL_VERSIONS)}; got {self.mode!r}{hint}"
|
|
)
|
|
|
|
self._folded_extensions = _fold_extensions(self.extensions)
|
|
|
|
srv = self.server
|
|
if isinstance(srv, MCPServer):
|
|
srv = srv._lowlevel_server # pyright: ignore[reportPrivateUsage]
|
|
if isinstance(srv, Server):
|
|
self._connect = _connect_inproc(srv)
|
|
elif isinstance(srv, str):
|
|
self._connect = _connect_transport(streamable_http_client(srv))
|
|
elif isinstance(srv, StdioServerParameters):
|
|
self._connect = _connect_transport(stdio_client(srv))
|
|
else:
|
|
self._connect = _connect_transport(srv)
|
|
|
|
if self.cache is not None:
|
|
config = self.cache
|
|
# Only the hash below leaves this scope - the raw identity may carry credentials; never log or store it.
|
|
target_id = config.target_id
|
|
if target_id is None and isinstance(self.server, str):
|
|
target_id = _strip_userinfo(self.server)
|
|
if target_id is None:
|
|
if config.store is not None:
|
|
raise ValueError(
|
|
"a custom cache store requires CacheConfig.target_id when the server is not a URL: "
|
|
"in-process servers and Transport instances get a random per-client identity, so "
|
|
"their entries in a shared store could never be served to another client"
|
|
)
|
|
target_id = uuid.uuid4().hex
|
|
self._response_cache = ClientResponseCache(
|
|
store=config.store if config.store is not None else InMemoryResponseCacheStore(),
|
|
partition=config.partition,
|
|
arm_id=hashlib.sha256(target_id.encode()).hexdigest(),
|
|
default_ttl_ms=config.default_ttl_ms,
|
|
clock=config.clock,
|
|
share_public=config.share_public,
|
|
# Lazy: the negotiated version is unknown until __aenter__'s handshake.
|
|
negotiated_version=lambda: self._session.protocol_version if self._session is not None else None,
|
|
)
|
|
|
|
async def _build_session(self, exit_stack: AsyncExitStack) -> ClientSession:
|
|
"""Enter the resolved connector and return an un-entered ClientSession."""
|
|
dispatcher = await self._connect(exit_stack, self.mode, self.raise_exceptions)
|
|
message_handler = self.message_handler
|
|
if self._response_cache is not None:
|
|
message_handler = _evicting_message_handler(self._response_cache, self.message_handler)
|
|
return ClientSession(
|
|
dispatcher=dispatcher,
|
|
read_timeout_seconds=self.read_timeout_seconds,
|
|
sampling_callback=self.sampling_callback,
|
|
sampling_capabilities=self.sampling_capabilities,
|
|
list_roots_callback=self.list_roots_callback,
|
|
logging_callback=self.logging_callback,
|
|
log_level=self.log_level,
|
|
message_handler=message_handler,
|
|
client_info=self.client_info,
|
|
elicitation_callback=self.elicitation_callback,
|
|
extensions=self._folded_extensions.ad,
|
|
result_claims=self._folded_extensions.claims,
|
|
notification_bindings=self._folded_extensions.bindings,
|
|
)
|
|
|
|
async def __aenter__(self) -> Client:
|
|
"""Enter the async context manager."""
|
|
if self._entered:
|
|
raise RuntimeError("Client is already entered; cannot reenter")
|
|
self._entered = True
|
|
|
|
async with AsyncExitStack() as exit_stack:
|
|
session = await self._build_session(exit_stack)
|
|
session = await exit_stack.enter_async_context(session)
|
|
|
|
if self.mode == "legacy":
|
|
await session.initialize()
|
|
elif self.mode == "auto":
|
|
await negotiate_auto(session)
|
|
else:
|
|
session.adopt(self.prior_discover or _synthesize_discover(self.mode))
|
|
|
|
# Only publish the session after the handshake succeeds, so `_session is not None`
|
|
# implies the protocol_version/server_capabilities are populated (server_info
|
|
# stays optional: 2026-era servers may not identify themselves). If the
|
|
# handshake raised above, the local exit_stack unwinds the transport for us.
|
|
self._session = session
|
|
self._exit_stack = exit_stack.pop_all()
|
|
return self
|
|
|
|
async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: Any) -> None:
|
|
"""Exit the async context manager."""
|
|
if self._exit_stack: # pragma: no branch
|
|
await self._exit_stack.__aexit__(exc_type, exc_val, exc_tb)
|
|
self._session = None
|
|
|
|
@property
|
|
def session(self) -> ClientSession:
|
|
"""Get the underlying ClientSession.
|
|
|
|
This provides access to the full ClientSession API for advanced use cases.
|
|
|
|
Raises:
|
|
RuntimeError: If accessed before entering the context manager.
|
|
"""
|
|
if self._session is None:
|
|
raise RuntimeError("Client must be used within an async context manager")
|
|
return self._session
|
|
|
|
# TODO(maxisbey): the by-construction shape is for __aenter__ to return a connected-view
|
|
# type whose protocol_version/server_capabilities are non-Optional fields,
|
|
# eliminating these guards (and the one in .session). Same family as resolving the
|
|
# transport/connector at __post_init__ so the Optional internal fields disappear.
|
|
# (server_info stays Optional even connected: the 2026-era stamp is optional.)
|
|
@property
|
|
def protocol_version(self) -> str:
|
|
"""Negotiated protocol version (set by initialize/discover/adopt during ``__aenter__``)."""
|
|
return _connected(self.session.protocol_version)
|
|
|
|
@property
|
|
def server_info(self) -> Implementation | None:
|
|
"""Server name/version, or `None` when the server did not identify itself.
|
|
|
|
Legacy connections always carry it (`InitializeResult.serverInfo` is
|
|
required); on 2026-era connections the `_meta` `serverInfo` stamp is
|
|
optional, so an anonymous server reads as `None`.
|
|
"""
|
|
return self.session.server_info
|
|
|
|
@property
|
|
def server_capabilities(self) -> ServerCapabilities:
|
|
"""Server capabilities (set by initialize/discover/adopt during ``__aenter__``)."""
|
|
return _connected(self.session.server_capabilities)
|
|
|
|
@property
|
|
def instructions(self) -> str | None:
|
|
"""Server-provided instructions text, if any."""
|
|
return self.session.instructions
|
|
|
|
@deprecated(
|
|
"ping is removed as of 2026-07-28; the method only works under mode='legacy'.",
|
|
category=MCPDeprecationWarning,
|
|
)
|
|
async def send_ping(self, *, meta: RequestParamsMeta | None = None) -> EmptyResult:
|
|
"""Send a ping request to the server."""
|
|
return await self.session.send_ping(meta=meta)
|
|
|
|
@deprecated(
|
|
"Client-to-server progress is deprecated as of 2026-07-28; progress is server-to-client only.",
|
|
category=MCPDeprecationWarning,
|
|
)
|
|
async def send_progress_notification(
|
|
self,
|
|
progress_token: str | int,
|
|
progress: float,
|
|
total: float | None = None,
|
|
message: str | None = None,
|
|
) -> None:
|
|
"""Send a progress notification to the server."""
|
|
await self.session.send_progress_notification( # pyright: ignore[reportDeprecated]
|
|
progress_token=progress_token,
|
|
progress=progress,
|
|
total=total,
|
|
message=message,
|
|
)
|
|
|
|
@deprecated("The logging capability is deprecated as of 2026-07-28 (SEP-2577).", category=MCPDeprecationWarning)
|
|
async def set_logging_level(self, level: LoggingLevel, *, meta: RequestParamsMeta | None = None) -> EmptyResult:
|
|
"""Set the logging level on the server."""
|
|
return await self.session.set_logging_level(level=level, meta=meta) # pyright: ignore[reportDeprecated]
|
|
|
|
async def _cached_fetch(
|
|
self,
|
|
method: str,
|
|
*,
|
|
cursor: str | None,
|
|
meta: RequestParamsMeta | None,
|
|
cache_mode: CacheMode,
|
|
send: Callable[[], Awaitable[_CacheableT]],
|
|
absorb: Callable[[_CacheableT], _CacheableT] | None = None,
|
|
) -> _CacheableT:
|
|
"""Serve one of the four list verbs through the response cache.
|
|
|
|
`absorb` (tools/list only) re-applies session-side derived state to a served cache hit.
|
|
"""
|
|
cache = self._response_cache
|
|
if cache is None or cache_mode == "bypass":
|
|
return await send()
|
|
# A closed (or never-entered) client must raise, never serve cached entries.
|
|
_ = self.session
|
|
if meta is not None and cache_mode == "use":
|
|
# meta (a progress token, tracing fields) expects a wire request; fetch and replace the entry.
|
|
cache_mode = "refresh"
|
|
if cursor is not None:
|
|
# Continuation pages skip the cache, but an expired cursor means the listing changed (spec SHOULD evict).
|
|
try:
|
|
return await send()
|
|
except MCPError as e:
|
|
if e.code == INVALID_PARAMS:
|
|
await cache.evict_method(method)
|
|
raise
|
|
if cache_mode == "use" and (hit := await cache.read(method, "")) is not None:
|
|
# The hit is a private deep copy, so absorption may mutate it freely.
|
|
served = cast(_CacheableT, hit)
|
|
return served if absorb is None else absorb(served)
|
|
gen = cache.capture(method, "")
|
|
result = await send()
|
|
await cache.write(method, "", result, gen, cache_mode)
|
|
return result
|
|
|
|
async def list_resources(
|
|
self,
|
|
*,
|
|
cursor: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
cache_mode: CacheMode = "use",
|
|
) -> ListResourcesResult:
|
|
"""List available resources from the server."""
|
|
return await self._cached_fetch(
|
|
"resources/list",
|
|
cursor=cursor,
|
|
meta=meta,
|
|
cache_mode=cache_mode,
|
|
send=lambda: self.session.list_resources(params=PaginatedRequestParams(cursor=cursor, _meta=meta)),
|
|
)
|
|
|
|
async def list_resource_templates(
|
|
self,
|
|
*,
|
|
cursor: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
cache_mode: CacheMode = "use",
|
|
) -> ListResourceTemplatesResult:
|
|
"""List available resource templates from the server."""
|
|
return await self._cached_fetch(
|
|
"resources/templates/list",
|
|
cursor=cursor,
|
|
meta=meta,
|
|
cache_mode=cache_mode,
|
|
send=lambda: self.session.list_resource_templates(params=PaginatedRequestParams(cursor=cursor, _meta=meta)),
|
|
)
|
|
|
|
async def read_resource(
|
|
self,
|
|
uri: str,
|
|
*,
|
|
input_responses: InputResponses | None = None,
|
|
request_state: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
cache_mode: CacheMode = "use",
|
|
) -> ReadResourceResult:
|
|
"""Read a resource from the server.
|
|
|
|
If the server returns an `InputRequiredResult`, the embedded input
|
|
requests are dispatched to this client's sampling / elicitation / roots
|
|
callbacks and the read is retried automatically (up to
|
|
`input_required_max_rounds`).
|
|
|
|
Args:
|
|
uri: The URI of the resource to read.
|
|
input_responses: Responses to seed the first call with (e.g. when
|
|
resuming from a persisted `InputRequiredResult`).
|
|
request_state: Opaque state to seed the first call with.
|
|
meta: Additional metadata for the request.
|
|
cache_mode: Cache behavior for this call (see `CacheMode`); seeded
|
|
calls (`input_responses` or `request_state` set) ignore it.
|
|
|
|
Returns:
|
|
The resource content.
|
|
|
|
Raises:
|
|
InputRequiredRoundsExceededError: `input_required_max_rounds` exhausted.
|
|
MCPError: A callback returned `ErrorData` for an embedded input request.
|
|
pydantic.ValidationError: The server returned a result that does not
|
|
conform to the negotiated protocol version.
|
|
"""
|
|
|
|
async def retry(r: InputResponses | None, s: str | None) -> ReadResourceResult | InputRequiredResult:
|
|
return await self.session.read_resource(
|
|
uri, input_responses=r, request_state=s, meta=meta, allow_input_required=True
|
|
)
|
|
|
|
# Seeded calls resume a specific exchange and must never be cached (spec MUST).
|
|
seeded = input_responses is not None or request_state is not None
|
|
cache = None if seeded else self._response_cache
|
|
if cache is None or cache_mode == "bypass":
|
|
return await self._drive_input_required(await retry(input_responses, request_state), retry)
|
|
# A closed (or never-entered) client must raise, never serve cached entries.
|
|
_ = self.session
|
|
if meta is not None and cache_mode == "use":
|
|
# Calls carrying meta always reach the server (mirrors `_cached_fetch`).
|
|
cache_mode = "refresh"
|
|
if cache_mode == "use" and (hit := await cache.read("resources/read", uri)) is not None:
|
|
# Only terminal first-round results are stored, so a hit legitimately skips the driver.
|
|
return cast(ReadResourceResult, hit)
|
|
gen = cache.capture("resources/read", uri)
|
|
first = await retry(None, None)
|
|
if not isinstance(first, InputRequiredResult):
|
|
await cache.write("resources/read", uri, first, gen, cache_mode)
|
|
elif cache_mode == "refresh":
|
|
# The refresh superseded whatever was cached, but an input_required resolution
|
|
# cannot be stored: purge the warm entry so it cannot be served again.
|
|
await cache.evict_key("resources/read", uri)
|
|
# Driver rounds carry inputResponses, so a terminal result reached through them is never cached (spec MUST).
|
|
return await self._drive_input_required(first, retry)
|
|
|
|
def listen(
|
|
self,
|
|
*,
|
|
tools_list_changed: bool = False,
|
|
prompts_list_changed: bool = False,
|
|
resources_list_changed: bool = False,
|
|
resource_subscriptions: Sequence[str] = (),
|
|
) -> AbstractAsyncContextManager[Subscription]:
|
|
"""Open a `subscriptions/listen` stream of typed change events (2026-07-28 only).
|
|
|
|
Keyword args mirror the wire `SubscriptionFilter`; entering waits for the ack (honored subset: `sub.honored`):
|
|
|
|
async with client.listen(tools_list_changed=True) as sub:
|
|
async for event in sub:
|
|
tools = await client.list_tools() # refetch on change
|
|
|
|
A graceful close ends the loop; an abrupt drop raises `SubscriptionLost`. No replay: re-listen and refetch.
|
|
|
|
Raises:
|
|
ListenNotSupportedError: The negotiated protocol version predates 2026-07-28.
|
|
MCPError: The server rejected the request or the connection failed first.
|
|
SubscriptionLost: The stream ended before it was acknowledged.
|
|
TimeoutError: The read timeout elapsed before the acknowledgment.
|
|
"""
|
|
return _listen(
|
|
self.session,
|
|
tools_list_changed=tools_list_changed,
|
|
prompts_list_changed=prompts_list_changed,
|
|
resources_list_changed=resources_list_changed,
|
|
resource_subscriptions=resource_subscriptions,
|
|
on_event=self._evict_for_listen_event if self._response_cache is not None else None,
|
|
)
|
|
|
|
async def _evict_for_listen_event(self, event: ServerEvent) -> None:
|
|
"""Finish response-cache eviction before a listen consumer can refetch.
|
|
|
|
Without it the iterator wakes first and refetches a still-warm entry, with no
|
|
corrective wake (events are deduplicated level triggers). The tee path repeats
|
|
the eviction; deliberate: idempotent, and it covers non-iterating consumers.
|
|
"""
|
|
cache = self._response_cache
|
|
assert cache is not None # installed as the event barrier only when a cache exists
|
|
try:
|
|
await cache.evict_for_notification(event_to_notification(event, {}))
|
|
except Exception: # boundary: eviction reaches user store code; a cache fault must not block delivery
|
|
logger.exception("Response cache eviction failed; the event is still delivered")
|
|
|
|
@deprecated(
|
|
"resources/subscribe is removed as of 2026-07-28; use Client.listen() instead.",
|
|
category=MCPDeprecationWarning,
|
|
)
|
|
async def subscribe_resource(self, uri: str, *, meta: RequestParamsMeta | None = None) -> EmptyResult:
|
|
"""Subscribe to resource updates (2025-era servers only)."""
|
|
return await self.session.subscribe_resource(uri, meta=meta) # pyright: ignore[reportDeprecated]
|
|
|
|
@deprecated(
|
|
"resources/unsubscribe is removed as of 2026-07-28; use Client.listen() instead.",
|
|
category=MCPDeprecationWarning,
|
|
)
|
|
async def unsubscribe_resource(self, uri: str, *, meta: RequestParamsMeta | None = None) -> EmptyResult:
|
|
"""Unsubscribe from resource updates (2025-era servers only)."""
|
|
return await self.session.unsubscribe_resource(uri, meta=meta) # pyright: ignore[reportDeprecated]
|
|
|
|
async def call_tool(
|
|
self,
|
|
name: str,
|
|
arguments: dict[str, Any] | None = None,
|
|
read_timeout_seconds: float | None = None,
|
|
progress_callback: ProgressFnT | None = None,
|
|
*,
|
|
input_responses: InputResponses | None = None,
|
|
request_state: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
) -> CallToolResult:
|
|
"""Call a tool on the server.
|
|
|
|
If the server returns an `InputRequiredResult`, the embedded input
|
|
requests are dispatched to this client's sampling / elicitation / roots
|
|
callbacks and the call is retried automatically (up to
|
|
`input_required_max_rounds`). To drive the loop yourself — e.g. to
|
|
persist `request_state` across process restarts — use
|
|
`client.session.call_tool(..., allow_input_required=True)`. Persisted
|
|
state is still subject to the server's TTL, request binding, and key
|
|
lifetime; a server on the default process-local key rejects it after a restart.
|
|
|
|
Result shapes claimed by this client's `extensions` are finished by the
|
|
owning claim's resolver, whose `CallToolResult` is returned; resolver
|
|
exceptions propagate as-is. To receive the claimed shape yourself, use
|
|
`client.session.call_tool(..., allow_claimed=True)`.
|
|
|
|
Args:
|
|
name: The name of the tool to call.
|
|
arguments: Arguments to pass to the tool.
|
|
read_timeout_seconds: Timeout for each underlying `tools/call` round.
|
|
progress_callback: Callback for progress updates.
|
|
input_responses: Responses to seed the first call with (e.g. when
|
|
resuming from a persisted `InputRequiredResult`).
|
|
request_state: Opaque state to seed the first call with.
|
|
meta: Additional metadata for the request.
|
|
|
|
Returns:
|
|
The tool result.
|
|
|
|
Raises:
|
|
InputRequiredRoundsExceededError: `input_required_max_rounds` exhausted.
|
|
MCPError: A callback returned `ErrorData` for an embedded input request.
|
|
pydantic.ValidationError: The server returned a result that does not
|
|
conform to the negotiated protocol version.
|
|
"""
|
|
|
|
async def retry(r: InputResponses | None, s: str | None) -> CallToolResult | InputRequiredResult | Result:
|
|
return await self.session.call_tool(
|
|
name,
|
|
arguments,
|
|
read_timeout_seconds=read_timeout_seconds,
|
|
progress_callback=progress_callback,
|
|
input_responses=r,
|
|
request_state=s,
|
|
meta=meta,
|
|
allow_input_required=True,
|
|
# Input rounds resolve before a claimed result, so a claim may end any round.
|
|
allow_claimed=True,
|
|
)
|
|
|
|
result = await self._drive_input_required(await retry(input_responses, request_state), retry)
|
|
if isinstance(result, CallToolResult):
|
|
return result
|
|
# Only claimed shapes reach this point, so the lookup is total.
|
|
claim = self._folded_extensions.by_model[type(result)]
|
|
final = await claim.resolve(
|
|
result,
|
|
ClaimContext(session=self.session, tool_name=name, read_timeout_seconds=read_timeout_seconds),
|
|
)
|
|
if not final.is_error:
|
|
# Match the direct path: revalidate the output schema, but never for isError results.
|
|
await self.session.validate_tool_result(name, final)
|
|
return final
|
|
|
|
async def list_prompts(
|
|
self,
|
|
*,
|
|
cursor: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
cache_mode: CacheMode = "use",
|
|
) -> ListPromptsResult:
|
|
"""List available prompts from the server."""
|
|
return await self._cached_fetch(
|
|
"prompts/list",
|
|
cursor=cursor,
|
|
meta=meta,
|
|
cache_mode=cache_mode,
|
|
send=lambda: self.session.list_prompts(params=PaginatedRequestParams(cursor=cursor, _meta=meta)),
|
|
)
|
|
|
|
async def get_prompt(
|
|
self,
|
|
name: str,
|
|
arguments: dict[str, str] | None = None,
|
|
*,
|
|
input_responses: InputResponses | None = None,
|
|
request_state: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
) -> GetPromptResult:
|
|
"""Get a prompt from the server.
|
|
|
|
If the server returns an `InputRequiredResult`, the embedded input
|
|
requests are dispatched to this client's sampling / elicitation / roots
|
|
callbacks and the get is retried automatically (up to
|
|
`input_required_max_rounds`).
|
|
|
|
Args:
|
|
name: The name of the prompt.
|
|
arguments: Arguments to pass to the prompt.
|
|
input_responses: Responses to seed the first call with (e.g. when
|
|
resuming from a persisted `InputRequiredResult`).
|
|
request_state: Opaque state to seed the first call with.
|
|
meta: Additional metadata for the request.
|
|
|
|
Returns:
|
|
The prompt content.
|
|
|
|
Raises:
|
|
InputRequiredRoundsExceededError: `input_required_max_rounds` exhausted.
|
|
MCPError: A callback returned `ErrorData` for an embedded input request.
|
|
pydantic.ValidationError: The server returned a result that does not
|
|
conform to the negotiated protocol version.
|
|
"""
|
|
|
|
async def retry(r: InputResponses | None, s: str | None) -> GetPromptResult | InputRequiredResult:
|
|
return await self.session.get_prompt(
|
|
name, arguments, input_responses=r, request_state=s, meta=meta, allow_input_required=True
|
|
)
|
|
|
|
return await self._drive_input_required(await retry(input_responses, request_state), retry)
|
|
|
|
async def _drive_input_required(
|
|
self,
|
|
first: _ResultT | InputRequiredResult,
|
|
retry: Callable[[InputResponses | None, str | None], Awaitable[_ResultT | InputRequiredResult]],
|
|
) -> _ResultT:
|
|
"""Hand an `InputRequiredResult` to the SEP-2322 driver, or pass a terminal result through.
|
|
|
|
`dispatch` routes each embedded request through the same callback table
|
|
that serves legacy server→client RPCs, so the two paths stay
|
|
behaviourally identical by construction.
|
|
"""
|
|
if not isinstance(first, InputRequiredResult):
|
|
return first
|
|
session = self.session
|
|
|
|
async def dispatch(key: str, req: InputRequest) -> InputResponse | ErrorData:
|
|
ctx = ClientRequestContext(session=session, request_id=key, meta=req.params.meta if req.params else None)
|
|
return await session.dispatch_input_request(ctx, req)
|
|
|
|
return await run_input_required_driver(
|
|
first, dispatch=dispatch, retry=retry, max_rounds=self.input_required_max_rounds
|
|
)
|
|
|
|
async def complete(
|
|
self,
|
|
ref: ResourceTemplateReference | PromptReference,
|
|
argument: dict[str, str],
|
|
context_arguments: dict[str, str] | None = None,
|
|
) -> CompleteResult:
|
|
"""Get completions for a prompt or resource template argument.
|
|
|
|
Args:
|
|
ref: Reference to the prompt or resource template
|
|
argument: The argument to complete
|
|
context_arguments: Additional context arguments
|
|
|
|
Returns:
|
|
Completion suggestions.
|
|
"""
|
|
return await self.session.complete(ref=ref, argument=argument, context_arguments=context_arguments)
|
|
|
|
async def list_tools(
|
|
self,
|
|
*,
|
|
cursor: str | None = None,
|
|
meta: RequestParamsMeta | None = None,
|
|
cache_mode: CacheMode = "use",
|
|
) -> ListToolsResult:
|
|
"""List available tools from the server."""
|
|
return await self._cached_fetch(
|
|
"tools/list",
|
|
cursor=cursor,
|
|
meta=meta,
|
|
cache_mode=cache_mode,
|
|
send=lambda: self.session.list_tools(params=PaginatedRequestParams(cursor=cursor, _meta=meta)),
|
|
# A cache hit skips session.list_tools, so the session re-absorbs the served
|
|
# listing to rebuild its derived per-tool state. Hits are cursorless, but a
|
|
# cached page 1 can carry next_cursor - never prune on a partial listing.
|
|
absorb=lambda hit: self.session._absorb_tool_listing( # pyright: ignore[reportPrivateUsage]
|
|
hit, complete=hit.next_cursor is None
|
|
),
|
|
)
|
|
|
|
@deprecated("The roots capability is deprecated as of 2026-07-28 (SEP-2577).", category=MCPDeprecationWarning)
|
|
async def send_roots_list_changed(self) -> None:
|
|
"""Send a notification that the roots list has changed."""
|
|
await self.session.send_roots_list_changed() # pyright: ignore[reportDeprecated]
|