diff --git a/docs/client/callbacks.md b/docs/client/callbacks.md index 6b4e934cf9..53f6e563ee 100644 --- a/docs/client/callbacks.md +++ b/docs/client/callbacks.md @@ -135,7 +135,7 @@ Two more. Neither declares anything. `logging_callback` receives every `notifications/message` a server sends, as `LoggingMessageNotificationParams` (`level`, `logger`, `data`). Protocol logging is itself deprecated by the 2026-07-28 spec (**[Logging](../handlers/logging.md)** has what to do instead), so this callback exists for the servers that still emit it. -`message_handler` is the catch-all: every server notification reaches it (as well as its specific callback), and on a stream-backed transport so does every transport-level `Exception`. The one pattern worth knowing is `if isinstance(message, Exception): raise message`, so a broken connection fails loudly instead of vanishing. +`message_handler` is the catch-all: every server notification the session surfaces reaches it (as well as its specific callback), and on a stream-backed transport so does every transport-level `Exception`. Two never do: `notifications/cancelled` is applied by the SDK rather than surfaced, and a subscription acknowledgment for a live `listen()` stream is consumed by that stream. Annotate the parameter with `IncomingMessage` (`ServerNotification | Exception`, exported from `mcp.client`). The one pattern worth knowing is `if isinstance(message, Exception): raise message`, so a broken connection fails loudly instead of vanishing. ## Recap diff --git a/docs/migration.md b/docs/migration.md index afe0953631..c2ffe1e6c9 100644 --- a/docs/migration.md +++ b/docs/migration.md @@ -1651,7 +1651,9 @@ Behavior changes: - **`send_notification` no longer takes `related_request_id`, and `send_request` no longer accepts `ServerMessageMetadata`.** No client transport ever serialized these hints; progress and response correlation via `progressToken` and the request id is unaffected. - **Client callbacks now receive `mcp.client.ClientRequestContext`** (its `request_id` is always populated); the `mcp.shared.context.RequestContext` generic is deleted. Annotations spelled `RequestContext[ClientSession, Any]` become `ClientRequestContext` (details in [`RequestContext` type parameters simplified](#requestcontext-type-parameters-simplified)). -`mcp.shared.session` is now a compatibility module: `ProgressFnT` is re-exported (its home is `mcp.shared.dispatcher`), and `RequestResponder` remains as a typing-only stub so `MessageHandlerFnT` annotations keep importing. `RequestResponder.respond()` no longer exists, and neither do the cancellation-tracking members (`cancel()`, the `cancelled` and `in_flight` properties, the `on_complete` constructor argument) or `BaseSession._in_flight`; inbound cancellation is handled by `JSONRPCDispatcher`. +- **`message_handler` no longer receives requests.** Server-initiated requests are answered by the typed callbacks (`sampling_callback`, `elicitation_callback`, `list_roots_callback`), so the handler's parameter is now `IncomingMessage = ServerNotification | Exception`, exported from `mcp.client`. Replace the hand-written v1 union `RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception` with `IncomingMessage`; `RequestResponder` is gone (below), so the old annotation no longer imports. + +The `mcp.shared.session` module is gone. `RequestResponder` is removed — `respond()`, the cancellation-tracking members (`cancel()`, the `cancelled` and `in_flight` properties, the `on_complete` constructor argument) and `BaseSession._in_flight` have no replacement; inbound cancellation is handled by `JSONRPCDispatcher`. `ProgressFnT` now lives only in `mcp.shared.dispatcher`, and `RequestId` in `mcp_types`. ### Experimental Tasks support removed diff --git a/docs/whats-new.md b/docs/whats-new.md index 3f1188f8fc..4ae0de2bcb 100644 --- a/docs/whats-new.md +++ b/docs/whats-new.md @@ -146,7 +146,7 @@ Each of these is a section in the **[Migration Guide](migration.md)**: * The **WebSocket transport**, both sides, and the `mcp[ws]` extra. It was never part of the MCP specification. * The **experimental Tasks** API (`mcp.*.experimental`). 2026-07-28 moves tasks out of the core protocol and into an official extension ([SEP-2663](https://github.com/modelcontextprotocol/modelcontextprotocol/pull/2663)), which this SDK does not implement yet. -* `mcp.types`, `mcp.shared.version`, and `mcp.shared.progress` as import paths. +* `mcp.types`, `mcp.shared.version`, `mcp.shared.progress`, and `mcp.shared.session` (with the `RequestResponder` stub v1 `message_handler` annotations imported) as import paths. * The deprecated `streamablehttp_client` spelling, and the `get_session_id` callback from `streamable_http_client` (which now yields exactly two streams). * `McpError`, renamed **`MCPError`** with a direct `(code, message, data)` constructor. * `MCPServer.get_context()`, `mount_path=`, and the lowlevel `Server`'s decorator methods, ContextVar, and handler dicts. diff --git a/examples/stories/standalone_get/client.py b/examples/stories/standalone_get/client.py index aaf870f0e7..d2054ca8df 100644 --- a/examples/stories/standalone_get/client.py +++ b/examples/stories/standalone_get/client.py @@ -3,7 +3,7 @@ import anyio import mcp_types as types -from mcp.client import Client +from mcp.client import Client, IncomingMessage from stories._harness import Target, run_client @@ -13,7 +13,7 @@ async def main(target: Target, *, mode: str = "auto") -> None: received: list[types.ResourceListChangedNotification] = [] seen = anyio.Event() - async def on_message(message: object) -> None: + async def on_message(message: IncomingMessage) -> None: if isinstance(message, types.ResourceListChangedNotification): received.append(message) seen.set() diff --git a/examples/stories/stickynotes/client.py b/examples/stories/stickynotes/client.py index 56ca10f551..a5b0e41dad 100644 --- a/examples/stories/stickynotes/client.py +++ b/examples/stories/stickynotes/client.py @@ -4,7 +4,7 @@ import mcp_types as types from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS -from mcp.client import Client, ClientRequestContext +from mcp.client import Client, ClientRequestContext, IncomingMessage from stories._harness import Target, run_client @@ -18,7 +18,7 @@ async def on_elicit(context: ClientRequestContext, params: types.ElicitRequestPa return types.ElicitResult(action="cancel") return types.ElicitResult(action="accept", content={"confirm": answer == "confirm"}) - async def on_message(message: object) -> None: + async def on_message(message: IncomingMessage) -> None: if isinstance(message, types.ResourceListChangedNotification): list_changed.set() diff --git a/src/mcp/client/__init__.py b/src/mcp/client/__init__.py index 21581749d0..d6b07045ce 100644 --- a/src/mcp/client/__init__.py +++ b/src/mcp/client/__init__.py @@ -20,7 +20,7 @@ UnexpectedClaimedResult, advertise, ) -from mcp.client.session import ClientSession +from mcp.client.session import ClientSession, IncomingMessage __all__ = [ "CacheConfig", @@ -32,6 +32,7 @@ "ClientExtension", "ClientRequestContext", "ClientSession", + "IncomingMessage", "InMemoryResponseCacheStore", "InputRequiredRoundsExceededError", "NotificationBinding", diff --git a/src/mcp/client/__main__.py b/src/mcp/client/__main__.py index 5fa3ce109b..60e3b02390 100644 --- a/src/mcp/client/__main__.py +++ b/src/mcp/client/__main__.py @@ -9,11 +9,10 @@ import mcp_types as types from mcp.client._transport import ReadStream, WriteStream -from mcp.client.session import ClientSession +from mcp.client.session import ClientSession, IncomingMessage from mcp.client.sse import sse_client from mcp.client.stdio import StdioServerParameters, stdio_client from mcp.shared.message import SessionMessage -from mcp.shared.session import RequestResponder if not sys.warnoptions: warnings.simplefilter("ignore") @@ -22,9 +21,7 @@ logger = logging.getLogger("client") -async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, -) -> None: +async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): logger.error("Error: %s", message) return diff --git a/src/mcp/client/client.py b/src/mcp/client/client.py index 9e26d40859..aaf7c83b0b 100644 --- a/src/mcp/client/client.py +++ b/src/mcp/client/client.py @@ -52,6 +52,7 @@ ClientRequestContext, ClientSession, ElicitationFnT, + IncomingMessage, ListRootsFnT, LoggingFnT, MessageHandlerFnT, @@ -68,7 +69,6 @@ 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.session import RequestResponder from mcp.shared.subscriptions import event_to_notification logger = logging.getLogger(__name__) @@ -155,9 +155,7 @@ def _strip_userinfo(url: str) -> str: 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: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def handler(message: IncomingMessage) -> None: if isinstance(message, types.ServerNotification): try: await cache.evict_for_notification(message) diff --git a/src/mcp/client/session.py b/src/mcp/client/session.py index 36cb78df5d..1ef4ddf3a3 100644 --- a/src/mcp/client/session.py +++ b/src/mcp/client/session.py @@ -6,7 +6,7 @@ from functools import reduce from operator import or_ from types import TracebackType -from typing import Annotated, Any, Final, Literal, Protocol, cast, overload +from typing import Annotated, Any, Final, Literal, Protocol, TypeAlias, cast, overload import anyio import anyio.abc @@ -53,7 +53,6 @@ ) from mcp.shared.jsonrpc_dispatcher import JSONRPCDispatcher, cancelled_request_id_from_params from mcp.shared.message import ClientMessageMetadata, SessionMessage -from mcp.shared.session import RequestResponder from mcp.shared.subscriptions import SUBSCRIPTION_ID_META_KEY, event_from_wire from mcp.shared.transport_context import TransportContext @@ -172,16 +171,20 @@ class LoggingFnT(Protocol): async def __call__(self, params: types.LoggingMessageNotificationParams) -> None: ... # pragma: no branch +IncomingMessage: TypeAlias = types.ServerNotification | Exception +"""What `message_handler` receives: the server notifications the session surfaces, plus transport-level exceptions. + +`notifications/cancelled` is applied by the dispatcher and never surfaced, and a +`notifications/subscriptions/acknowledged` for a live `listen()` stream is consumed by that +stream, so neither reaches the handler. +""" + + class MessageHandlerFnT(Protocol): - async def __call__( - self, - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: ... # pragma: no branch + async def __call__(self, message: IncomingMessage) -> None: ... # pragma: no branch -async def _default_message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, -) -> None: +async def _default_message_handler(message: IncomingMessage) -> None: await anyio.lowlevel.checkpoint() @@ -331,8 +334,10 @@ class ClientSession: `dispatcher=`), enter as an async context manager, then call `initialize()`. The dispatcher owns the receive loop and request correlation; this class owns the typed MCP layer and the constructor - callbacks. Transport `Exception` items reach `message_handler` only when - the session builds its own dispatcher from a stream pair. + callbacks. Transport `Exception` items reach `message_handler` on any + stream-backed dispatcher (`JSONRPCDispatcher`), whether built here from a + stream pair or supplied without a stream-exception hook of its own; an + in-process `DirectDispatcher` carries none. Extension `result_claims` fold into tools/call parsing at `adopt()`; `notification_bindings` observe vendor notifications via bounded FIFOs. diff --git a/src/mcp/client/session_group.py b/src/mcp/client/session_group.py index 5f26a43365..a544cecbe8 100644 --- a/src/mcp/client/session_group.py +++ b/src/mcp/client/session_group.py @@ -25,8 +25,8 @@ from mcp.client.stdio import StdioServerParameters from mcp.client.streamable_http import streamable_http_client from mcp.shared._httpx_utils import create_mcp_http_client +from mcp.shared.dispatcher import ProgressFnT from mcp.shared.exceptions import MCPError -from mcp.shared.session import ProgressFnT class SseServerParameters(BaseModel): diff --git a/src/mcp/shared/session.py b/src/mcp/shared/session.py deleted file mode 100644 index f8f0a6d416..0000000000 --- a/src/mcp/shared/session.py +++ /dev/null @@ -1,22 +0,0 @@ -"""Compatibility names that outlived the removed v1 session layer (`BaseSession`).""" - -from typing import Generic, TypeVar - -from mcp_types import RequestParamsMeta - -from mcp.shared.dispatcher import ProgressFnT as ProgressFnT -from mcp.shared.message import MessageMetadata - -RequestId = str | int - -ReceiveRequestT = TypeVar("ReceiveRequestT") -SendResultT = TypeVar("SendResultT") - - -class RequestResponder(Generic[ReceiveRequestT, SendResultT]): - """Typing stub for the v1 responder; the SDK never instantiates it.""" - - request_id: RequestId - request_meta: RequestParamsMeta | None - request: ReceiveRequestT - message_metadata: MessageMetadata diff --git a/tests/client/test_client_caching.py b/tests/client/test_client_caching.py index aa6ed22e20..30f752eb6c 100644 --- a/tests/client/test_client_caching.py +++ b/tests/client/test_client_caching.py @@ -42,7 +42,7 @@ ) from mcp_types.version import LATEST_MODERN_VERSION -from mcp.client import Client +from mcp.client import Client, IncomingMessage from mcp.client._transport import TransportStreams from mcp.client.caching import ( CacheConfig, @@ -57,13 +57,10 @@ from mcp.shared.exceptions import MCPError from mcp.shared.memory import MessageStream, create_client_server_memory_streams from mcp.shared.message import SessionMessage -from mcp.shared.session import RequestResponder from tests.interaction._connect import BASE_URL, mounted_app pytestmark = pytest.mark.anyio -IncomingMessage = RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception - def _coordinator(client: Client) -> ClientResponseCache: cache = client._response_cache diff --git a/tests/client/test_logging_callback.py b/tests/client/test_logging_callback.py index d62b7e19b3..7ccdae3530 100644 --- a/tests/client/test_logging_callback.py +++ b/tests/client/test_logging_callback.py @@ -1,6 +1,5 @@ from typing import Literal -import mcp_types as types import pytest from mcp_types import ( LoggingMessageNotificationParams, @@ -8,8 +7,8 @@ ) from mcp import Client +from mcp.client import IncomingMessage from mcp.server.mcpserver import Context, MCPServer -from mcp.shared.session import RequestResponder class LoggingCollector: @@ -55,9 +54,7 @@ async def test_tool_with_log_dict( return True # Create a message handler to catch exceptions - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): # pragma: no cover raise message diff --git a/tests/client/test_notification_response.py b/tests/client/test_notification_response.py index 6724dfaf1b..b21e734fa3 100644 --- a/tests/client/test_notification_response.py +++ b/tests/client/test_notification_response.py @@ -16,8 +16,8 @@ from starlette.routing import Route from mcp import ClientSession, MCPError +from mcp.client import IncomingMessage from mcp.client.streamable_http import streamable_http_client -from mcp.shared.session import RequestResponder pytestmark = pytest.mark.anyio @@ -82,9 +82,7 @@ async def test_non_compliant_notification_response() -> None: """ returned_exception = None - async def message_handler( # pragma: no cover - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: # pragma: no cover nonlocal returned_exception if isinstance(message, Exception): returned_exception = message diff --git a/tests/client/test_session.py b/tests/client/test_session.py index 9c935ef18a..12fab4c8c9 100644 --- a/tests/client/test_session.py +++ b/tests/client/test_session.py @@ -37,7 +37,7 @@ from pydantic import FileUrl, ValidationError from mcp import MCPError -from mcp.client import ClientRequestContext +from mcp.client import ClientRequestContext, IncomingMessage from mcp.client.client import Client from mcp.client.session import DEFAULT_CLIENT_INFO, ClientSession from mcp.client.subscriptions import ToolsListChanged, listen @@ -45,7 +45,6 @@ from mcp.shared.direct_dispatcher import create_direct_dispatcher_pair from mcp.shared.dispatcher import CallOptions, DispatchContext, OnNotify, OnNotifyIntercept, OnRequest from mcp.shared.message import SessionMessage -from mcp.shared.session import RequestResponder from mcp.shared.subscriptions import SUBSCRIPTION_ID_META_KEY from mcp.shared.transport_context import TransportContext @@ -123,9 +122,7 @@ async def mock_server(): ) # Create a message handler to catch exceptions - async def message_handler( # pragma: no cover - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: # pragma: no cover if isinstance(message, Exception): raise message @@ -1228,10 +1225,8 @@ async def test_raising_notification_callbacks_over_direct_dispatch_cost_only_tha async def logging_callback(params: types.LoggingMessageNotificationParams) -> None: raise ValueError("logging callback boom") - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: - assert not isinstance(message, RequestResponder | Exception) + async def message_handler(message: IncomingMessage) -> None: + assert not isinstance(message, Exception) teed.append(message) raise ValueError("message handler boom") diff --git a/tests/interaction/README.md b/tests/interaction/README.md index 89be3d3abf..3060a240c1 100644 --- a/tests/interaction/README.md +++ b/tests/interaction/README.md @@ -40,7 +40,7 @@ flows — with a single subprocess test for stdio. ```text tests/interaction/ _requirements.py the requirements manifest (see below) - _helpers.py shared type aliases + the wire-recording transport + _helpers.py the wire-recording transport _connect.py the transport-parametrized connection factories conftest.py the connect fixture (the transport matrix) test_coverage.py enforces the manifest ↔ test contract diff --git a/tests/interaction/_helpers.py b/tests/interaction/_helpers.py index 0641aeab97..b335def7d0 100644 --- a/tests/interaction/_helpers.py +++ b/tests/interaction/_helpers.py @@ -1,28 +1,16 @@ """Shared helpers for the interaction suite. -Keep this module small: it exists only for (a) types that every test would otherwise have to -assemble from the SDK's internals to annotate a client callback, and (b) the recording transport -used by the wire-level tests. Server fixtures and assertion helpers belong in the test that uses -them. +Keep this module small: it exists only for the recording transport used by the wire-level +tests. Server fixtures and assertion helpers belong in the test that uses them. """ from types import TracebackType import anyio -from mcp_types import ClientResult, ServerNotification, ServerRequest from typing_extensions import Self from mcp.client._transport import ReadStream, Transport, TransportStreams, WriteStream from mcp.shared.message import SessionMessage -from mcp.shared.session import RequestResponder - -# TODO: this union is the parameter type of every client message handler (MessageHandlerFnT), -# but the SDK does not export a name for it -- writing a correctly-typed handler requires -# importing RequestResponder from mcp.shared.session and assembling the union by hand. It -# should be a named, exported alias next to MessageHandlerFnT (like ClientRequestContext is -# for the request callbacks), at which point this alias can be deleted. -IncomingMessage = RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception -"""Everything a client message handler can receive.""" class _RecordingReadStream: diff --git a/tests/interaction/lowlevel/test_cancellation.py b/tests/interaction/lowlevel/test_cancellation.py index 8361865db0..ecbf1a088d 100644 --- a/tests/interaction/lowlevel/test_cancellation.py +++ b/tests/interaction/lowlevel/test_cancellation.py @@ -29,13 +29,12 @@ ) from mcp import MCPError -from mcp.client import ClientRequestContext, ClientSession +from mcp.client import ClientRequestContext, ClientSession, IncomingMessage from mcp.server import Server, ServerRequestContext from mcp.shared.memory import MessageStream, create_client_server_memory_streams from mcp.shared.message import SessionMessage from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/lowlevel/test_elicitation.py b/tests/interaction/lowlevel/test_elicitation.py index b8393dd316..7024ef9e0f 100644 --- a/tests/interaction/lowlevel/test_elicitation.py +++ b/tests/interaction/lowlevel/test_elicitation.py @@ -28,12 +28,11 @@ ) from mcp import MCPError, UrlElicitationRequiredError -from mcp.client import ClientRequestContext, ClientSession +from mcp.client import ClientRequestContext, ClientSession, IncomingMessage from mcp.server import Server, ServerRequestContext from mcp.shared.memory import MessageStream, create_client_server_memory_streams from mcp.shared.message import SessionMessage from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/lowlevel/test_flows.py b/tests/interaction/lowlevel/test_flows.py index 9a6db1a45a..78ca716021 100644 --- a/tests/interaction/lowlevel/test_flows.py +++ b/tests/interaction/lowlevel/test_flows.py @@ -29,12 +29,11 @@ ) from mcp import MCPError, UrlElicitationRequiredError -from mcp.client import ClientRequestContext +from mcp.client import ClientRequestContext, IncomingMessage from mcp.server import Server, ServerRequestContext from mcp.server.session import ServerSession from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/lowlevel/test_list_changed.py b/tests/interaction/lowlevel/test_list_changed.py index e7d497ba2d..d978c0db51 100644 --- a/tests/interaction/lowlevel/test_list_changed.py +++ b/tests/interaction/lowlevel/test_list_changed.py @@ -26,9 +26,9 @@ ToolListChangedNotification, ) +from mcp.client import IncomingMessage from mcp.server import Server, ServerRequestContext from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/lowlevel/test_progress.py b/tests/interaction/lowlevel/test_progress.py index 025dded156..4d1f0d42af 100644 --- a/tests/interaction/lowlevel/test_progress.py +++ b/tests/interaction/lowlevel/test_progress.py @@ -15,12 +15,12 @@ from inline_snapshot import snapshot from mcp_types import CallToolResult, ProgressNotification, ProgressNotificationParams, ProgressToken, TextContent +from mcp.client import IncomingMessage from mcp.server import Server, ServerRequestContext from mcp.server.session import ServerSession -from mcp.shared.session import ProgressFnT +from mcp.shared.dispatcher import ProgressFnT from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/lowlevel/test_resources.py b/tests/interaction/lowlevel/test_resources.py index e911c65876..9e18393707 100644 --- a/tests/interaction/lowlevel/test_resources.py +++ b/tests/interaction/lowlevel/test_resources.py @@ -26,10 +26,10 @@ ) from mcp import MCPError +from mcp.client import IncomingMessage from mcp.server import Server, ServerRequestContext from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/mcpserver/test_context.py b/tests/interaction/mcpserver/test_context.py index 24728e2da8..0f07a128b7 100644 --- a/tests/interaction/mcpserver/test_context.py +++ b/tests/interaction/mcpserver/test_context.py @@ -17,12 +17,11 @@ from pydantic import BaseModel from mcp import MCPError -from mcp.client import ClientRequestContext +from mcp.client import ClientRequestContext, IncomingMessage from mcp.server.elicitation import AcceptedElicitation from mcp.server.mcpserver import Context, MCPServer from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/mcpserver/test_tools.py b/tests/interaction/mcpserver/test_tools.py index ad5db520b9..21a7163360 100644 --- a/tests/interaction/mcpserver/test_tools.py +++ b/tests/interaction/mcpserver/test_tools.py @@ -17,12 +17,12 @@ from pydantic import BaseModel, Field from mcp import MCPError +from mcp.client import IncomingMessage from mcp.server.mcpserver import Context, MCPServer from mcp.server.mcpserver.exceptions import ToolError from mcp.shared.exceptions import UrlElicitationRequiredError from tests._stamp import Unstamp from tests.interaction._connect import Connect -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/interaction/transports/test_streamable_http.py b/tests/interaction/transports/test_streamable_http.py index bb6dec5695..cb22e7ab87 100644 --- a/tests/interaction/transports/test_streamable_http.py +++ b/tests/interaction/transports/test_streamable_http.py @@ -23,12 +23,11 @@ ) from pydantic import BaseModel -from mcp.client import ClientRequestContext +from mcp.client import ClientRequestContext, IncomingMessage from mcp.server.elicitation import AcceptedElicitation from mcp.server.mcpserver import Context, MCPServer from mcp.shared.exceptions import MCPError from tests.interaction._connect import connect_over_streamable_http -from tests.interaction._helpers import IncomingMessage from tests.interaction._requirements import requirement pytestmark = pytest.mark.anyio diff --git a/tests/server/mcpserver/test_integration.py b/tests/server/mcpserver/test_integration.py index f6361f7574..2c57aab086 100644 --- a/tests/server/mcpserver/test_integration.py +++ b/tests/server/mcpserver/test_integration.py @@ -14,7 +14,6 @@ import pytest from inline_snapshot import snapshot from mcp_types import ( - ClientResult, CreateMessageRequestParams, CreateMessageResult, ElicitRequestParams, @@ -30,7 +29,6 @@ ResourceListChangedNotification, ResourceTemplateReference, ServerNotification, - ServerRequest, TextContent, TextResourceContents, ToolListChangedNotification, @@ -48,8 +46,7 @@ structured_output, tool_progress, ) -from mcp.client import Client, ClientRequestContext -from mcp.shared.session import RequestResponder +from mcp.client import Client, ClientRequestContext, IncomingMessage pytestmark = pytest.mark.anyio @@ -63,9 +60,7 @@ def __init__(self): self.resource_notifications: list[NotificationParams | None] = [] self.tool_notifications: list[NotificationParams | None] = [] - async def handle_generic_notification( - self, message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception - ) -> None: + async def handle_generic_notification(self, message: IncomingMessage) -> None: """Handle any server notification and route to appropriate handler.""" if isinstance(message, ServerNotification): # pragma: no branch if isinstance(message, ProgressNotification): @@ -180,7 +175,7 @@ async def test_tool_progress() -> None: """Test tool progress reporting.""" collector = NotificationCollector() - async def message_handler(message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception): + async def message_handler(message: IncomingMessage): await collector.handle_generic_notification(message) if isinstance(message, Exception): # pragma: no cover raise message @@ -259,7 +254,7 @@ async def test_notifications() -> None: """Test notifications and logging functionality.""" collector = NotificationCollector() - async def message_handler(message: RequestResponder[ServerRequest, ClientResult] | ServerNotification | Exception): + async def message_handler(message: IncomingMessage): await collector.handle_generic_notification(message) if isinstance(message, Exception): # pragma: no cover raise message diff --git a/tests/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index d7eeccdfdb..aeef25a278 100644 --- a/tests/shared/test_streamable_http.py +++ b/tests/shared/test_streamable_http.py @@ -43,7 +43,7 @@ from starlette.types import Message, Scope from mcp import MCPError -from mcp.client import ClientRequestContext +from mcp.client import ClientRequestContext, IncomingMessage from mcp.client.session import ClientSession from mcp.client.streamable_http import StreamableHTTPTransport, streamable_http_client from mcp.server import Server, ServerRequestContext @@ -64,7 +64,6 @@ from mcp.shared._compat import resync_tracer from mcp.shared._context_streams import create_context_streams from mcp.shared.message import ClientMessageMetadata, ServerMessageMetadata, SessionMessage -from mcp.shared.session import RequestResponder from tests.interaction.transports import StreamingASGITransport # Test constants @@ -968,9 +967,7 @@ async def test_streamable_http_client_get_stream(basic_app: Starlette) -> None: notifications_received: list[types.ServerNotification] = [] # Define message handler to capture notifications - async def message_handler( # pragma: no branch - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: # pragma: no branch if isinstance(message, types.ServerNotification): # pragma: no branch notifications_received.append(message) @@ -1128,9 +1125,7 @@ async def test_streamable_http_client_resumption(event_app: tuple[SimpleEventSto first_notification_received = anyio.Event() resumption_token_received = anyio.Event() - async def message_handler( # pragma: no branch - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: # pragma: no branch if isinstance(message, types.ServerNotification): # pragma: no branch captured_notifications.append(message) # Look for our first notification @@ -1798,14 +1793,11 @@ async def test_streamable_http_client_auto_reconnects( _, app = event_app captured_notifications: list[str] = [] - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): # pragma: no branch return # pragma: no cover - if isinstance(message, types.ServerNotification): # pragma: no branch - if isinstance(message, types.LoggingMessageNotification): # pragma: no branch - captured_notifications.append(str(message.params.data)) + if isinstance(message, types.LoggingMessageNotification): # pragma: no branch + captured_notifications.append(str(message.params.data)) async with ( make_client(app) as http_client, @@ -1866,14 +1858,11 @@ async def test_streamable_http_sse_polling_full_cycle( _, app = event_app all_notifications: list[str] = [] - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): # pragma: no branch return # pragma: no cover - if isinstance(message, types.ServerNotification): # pragma: no branch - if isinstance(message, types.LoggingMessageNotification): # pragma: no branch - all_notifications.append(str(message.params.data)) + if isinstance(message, types.LoggingMessageNotification): # pragma: no branch + all_notifications.append(str(message.params.data)) async with ( make_client(app) as http_client, @@ -1909,14 +1898,11 @@ async def test_streamable_http_events_replayed_after_disconnect( _, app = event_app notification_data: list[str] = [] - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): # pragma: no branch return # pragma: no cover - if isinstance(message, types.ServerNotification): # pragma: no branch - if isinstance(message, types.LoggingMessageNotification): # pragma: no branch - notification_data.append(str(message.params.data)) + if isinstance(message, types.LoggingMessageNotification): # pragma: no branch + notification_data.append(str(message.params.data)) async with ( make_client(app) as http_client, @@ -2042,14 +2028,11 @@ async def test_standalone_get_stream_reconnection(event_app: tuple[SimpleEventSt _, app = event_app received_notifications: list[str] = [] - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: if isinstance(message, Exception): return # pragma: no cover - if isinstance(message, types.ServerNotification): # pragma: no branch - if isinstance(message, types.ResourceUpdatedNotification): # pragma: no branch - received_notifications.append(str(message.params.uri)) + if isinstance(message, types.ResourceUpdatedNotification): # pragma: no branch + received_notifications.append(str(message.params.uri)) async with ( make_client(app) as http_client, @@ -2183,9 +2166,7 @@ async def test_standalone_stream_teardown_mid_listen_is_not_an_error(caplog: pyt app = Starlette(routes=[Mount("/mcp", app=session_manager.handle_request)]) notified = anyio.Event() - async def message_handler( - message: RequestResponder[types.ServerRequest, types.ClientResult] | types.ServerNotification | Exception, - ) -> None: + async def message_handler(message: IncomingMessage) -> None: # Only the standalone-stream notification is teed to the handler here. assert isinstance(message, types.ResourceUpdatedNotification) notified.set()