diff --git a/docs/migration.md b/docs/migration.md index 11569680e0..f987f20de3 100644 --- a/docs/migration.md +++ b/docs/migration.md @@ -863,7 +863,7 @@ Lifespans that set up process-wide state (connection pools, caches, background t ### Streamable HTTP: session manager, `EventStore`, and stateless mode unchanged -Beyond the constructor parameters that moved to `run()`/`streamable_http_app()` and the lifespan change above, the server-side Streamable HTTP machinery is as in v1: +Beyond the constructor parameters that moved to `run()`/`streamable_http_app()`, the lifespan change above, and the transport rework in the next section (which does not touch the public surface), the server-side Streamable HTTP API is as in v1: - `mcp.server.streamable_http` still exports the `EventStore` ABC (`store_event()`, `replay_events_after()`), `EventMessage`, `EventCallback`, `EventId`, and `StreamId` with unchanged signatures; a custom `EventStore` keeps importing `JSONRPCMessage` from `mcp.types`, unchanged. - `StreamableHTTPSessionManager` keeps its constructor and its `run()` / `handle_request()` methods (see [Lowlevel `Server`: what did not change](#lowlevel-server-what-did-not-change)); its `stateless=` parameter is unrelated to the removed [`Server.run(stateless=)` flag](#serverrun-no-longer-takes-a-stateless-flag). @@ -872,6 +872,59 @@ Beyond the constructor parameters that moved to `run()`/`streamable_http_app()` Only private attributes moved: `mcp._mcp_server` is now `mcp._lowlevel_server` (see [Registering lowlevel handlers from `MCPServer`](#registering-lowlevel-handlers-from-mcpserver)), and `_session_manager` now lives on that lowlevel `Server`. Prefer the public `mcp.session_manager` property to either. +### Streamable HTTP: `StreamableHTTPServerTransport` is driven per request, not per stream + +`StreamableHTTPServerTransport` no longer exposes a `connect()` context manager yielding a +`(read_stream, write_stream)` pair for you to run a server loop over. Each HTTP request is now +dispatched to the server's handlers directly: a request's outbound messages ride that request's +own response stream (backed by the optional `EventStore` for `Last-Event-ID` resumability), the +client's POSTed answers to server-initiated requests are correlated back by request id, and the +standalone GET stream is a further per-connection channel. The transport is the per-session +core; `StreamableHTTPSessionManager` binds one to the `Server` for each session (via the new +keyword-only `app` / `lifespan_state` constructor arguments) and routes requests to it. + +Nothing changes if you serve through `streamable_http_app()` / `run(transport="streamable-http")` +or mount `StreamableHTTPSessionManager` — the wire behaviour (session ids, GET stream, event +store, `ctx.close_sse_stream()`, `related_request_id` routing) is unchanged. Only code that +constructed a transport and consumed `transport.connect()` by hand needs to move to the session +manager: + +```python +# Before (v1) +transport = StreamableHTTPServerTransport(mcp_session_id=session_id, ...) +async with transport.connect() as (read_stream, write_stream): + await server.run(read_stream, write_stream, server.create_initialization_options()) + +# After (v2): serve the app the SDK builds ... +app = server.streamable_http_app(event_store=..., json_response=...) +``` + +... or, when composing your own Starlette/FastAPI app, mount a `StreamableHTTPSessionManager` and +enter `session_manager.run()` in the lifespan — see [Mounting the ASGI app](run/asgi.md) for the +full wiring. + +Behaviour clarified in the same change: + +- In JSON-response mode a *request-scoped* server-to-client request (`ctx.elicit()`, or any + `ctx.session` call carrying `related_request_id`) now raises `NoBackChannelError` — the POST's + single JSON body has no stream to carry the nested request, and previously the call would hang + waiting for an answer that could never be delivered. Connection-scoped sends (calls without + `related_request_id`) are unchanged and still ride the standalone GET stream. +- A GET carrying `Last-Event-ID` on a server without an `EventStore` opens the standalone stream + as a plain GET would, since there is nothing to replay. +- Two concurrent POSTs that share a JSON-RPC request id each keep their own response stream; the + second no longer silently takes over the first's queue. +- Stream ids handed to your `EventStore` are minted by the transport in its own session-scoped + namespace (previously the raw `str(request_id)` and a single global GET-stream key), so two + sessions sharing one store no longer collide, and a `Last-Event-ID` replay only releases frames + of the requesting session's own streams. Treat the ids as opaque. +- A failing `EventStore.store_event` degrades resumability for that message rather than taking + the stream down: the message is still delivered live (with no event id to resume from) and the + store's exception is logged, never sent to the client. +- A server-to-client request that can reach no client at all (no attached stream and nothing + storing it, or a request-scoped one in JSON-response mode) fails the calling handler with + `CONNECTION_CLOSED` instead of parking it for an answer that cannot arrive. + ### `MCPServer.get_context()` removed `MCPServer.get_context()` has been removed. Context is now injected by the framework and passed explicitly — there is no ambient ContextVar to read from. diff --git a/src/mcp/server/lowlevel/server.py b/src/mcp/server/lowlevel/server.py index 1cbd3f2bd6..3bcb301040 100644 --- a/src/mcp/server/lowlevel/server.py +++ b/src/mcp/server/lowlevel/server.py @@ -704,8 +704,9 @@ async def run( Thin wrapper over `serve_dual_era_loop`: enters the server lifespan, then drives the loop, serving the legacy handshake era and the modern per-request-envelope era (the client's first request decides which). - Transports with their own lifespan owner (the streamable-HTTP manager) - call `serve_loop` directly instead. + Transports with their own lifespan owner call `serve_loop` directly + instead (or, without a stream pair - the streamable-HTTP manager - + dispatch each request themselves). """ async with self.lifespan(self) as lifespan_context: await serve_dual_era_loop( diff --git a/src/mcp/server/runner.py b/src/mcp/server/runner.py index 6f9f7a8f74..c1b10f7ab1 100644 --- a/src/mcp/server/runner.py +++ b/src/mcp/server/runner.py @@ -230,7 +230,12 @@ async def _inner(ctx: ServerRequestContext[LifespanT, Any]) -> HandlerResult: result = _dump_result(await call(ctx)) if method == "initialize": # Commit only on chain success, so a middleware veto leaves no state. - # Race-free: the read loop is parked until this call returns. + # Race-free for the session's first handshake: the transport runs no + # other request until it returns (a stream driver's read loop is + # parked here; streamable HTTP holds later requests behind the + # in-progress initialize, and the session id only ships with its + # response). A repeated initialize on an established session (a + # recorded divergence) recommits alongside whatever is running. # TODO: this re-reads the wire `params`, so a middleware that rewrote # `ctx.params` (or `ctx.method`, or short-circuited without `call_next`) # can leave `connection.protocol_version` out of step with the @@ -477,9 +482,9 @@ async def serve_loop( """Drive ``server`` in handshake-only loop mode over a stream pair until the channel closes. Builds the loop-mode `JSONRPCDispatcher` + `Connection` and hands them to - `serve_connection`. The streamable-HTTP manager (which owns its lifespan - and serves the modern era on the single-exchange entry instead) calls - this; `Server.run` drives `serve_dual_era_loop`, which extends the same + `serve_connection`. For a transport that supplies a duplex message stream + pair but owns its own lifespan (so `Server.run`'s lifespan entry is not + wanted); `Server.run` drives `serve_dual_era_loop`, which extends the same dispatcher recipe (notably the `inline_methods={"initialize"}` rule) with era routing. """ diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 324fc3e04b..3caf694bf8 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -1,22 +1,36 @@ """StreamableHTTP Server Transport Module -This module implements an HTTP transport layer with Streamable HTTP. - -The transport handles bidirectional communication using HTTP requests and -responses, with streaming support for long-running operations. +This module implements the (2025-era, sessionful) Streamable HTTP transport. + +Each HTTP request is served directly: a POSTed JSON-RPC request is +dispatched to the server's handler kernel and its outbound messages flow into +a per-request `_MessageChannel` - the response's own SSE stream, backed by +the optional `EventStore` for resumability - rather than through a shared +message pipe. A POSTed JSON-RPC response resolves the server-to-client request +awaiting it; a POSTed notification is handled after the `202`. The standalone +GET stream is one more channel, connection-scoped, for messages related to +no request. + +`StreamableHTTPServerTransport` is therefore the per-session core (session +id, connection state, correlation of server-to-client requests, the open +channels); `StreamableHTTPSessionManager` creates one per `Mcp-Session-Id` +(or a fresh one per request in stateless mode) and routes ASGI requests to it. """ +from __future__ import annotations + import logging import re from abc import ABC, abstractmethod -from collections.abc import AsyncGenerator, Awaitable, Callable -from contextlib import asynccontextmanager -from dataclasses import dataclass +from collections.abc import Awaitable, Callable, Coroutine, Mapping +from dataclasses import dataclass, field from functools import partial from http import HTTPStatus -from typing import Any, Final +from typing import TYPE_CHECKING, Any, Final +from uuid import uuid4 import anyio +import anyio.abc import pydantic_core from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from mcp_types import ( @@ -28,6 +42,7 @@ ErrorData, JSONRPCError, JSONRPCMessage, + JSONRPCNotification, JSONRPCRequest, JSONRPCResponse, RequestId, @@ -40,11 +55,19 @@ from starlette.responses import Response from starlette.types import Receive, Scope, Send +from mcp.server.connection import Connection +from mcp.server.runner import ServerRunner, aclose_shielded from mcp.server.transport_security import TransportSecurityMiddleware, TransportSecuritySettings -from mcp.shared._context_streams import ContextReceiveStream, ContextSendStream, create_context_streams -from mcp.shared._stream_protocols import ReadStream, WriteStream +from mcp.shared._correlation import RequestCorrelator +from mcp.shared.dispatcher import CallOptions +from mcp.shared.exceptions import NoBackChannelError from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER -from mcp.shared.message import ServerMessageMetadata, SessionMessage +from mcp.shared.jsonrpc_dispatcher import cancelled_request_id_from_params, progress_token_from_params +from mcp.shared.message import ServerMessageMetadata +from mcp.shared.transport_context import TransportContext + +if TYPE_CHECKING: + from mcp.server.lowlevel.server import Server logger = logging.getLogger(__name__) @@ -60,16 +83,15 @@ # Special key for the standalone GET stream GET_STREAM_KEY = "_GET_stream" -# Buffer for the per-request `_request_streams` so the serial `message_router` -# can deposit a response and move on instead of head-of-line blocking the -# whole session on a lazily-started `sse_writer`. See #1764. +# Buffer between a channel and the SSE response draining it, so a handler can +# run this far ahead of a slow client before its own writes apply backpressure. REQUEST_STREAM_BUFFER_SIZE: Final = 16 # Error code answering a request that settled without a response (e.g. it was # cancelled) on this 2025-era wire, which ends a request's stream only with a # response. Mirrors LSP's RequestCancelled; not sent by the 2026 transports, where # the spec forbids answering a cancelled request. See -# `StreamableHTTPServerTransport._terminate_unanswered_request`. +# `StreamableHTTPServerTransport._settle_unanswered_request`. REQUEST_CANCELLED: Final = -32800 # Session ID validation pattern (visible ASCII characters ranging from 0x21 to 0x7E) @@ -141,23 +163,260 @@ async def replay_events_after( send_callback: A callback function to send events to the client Returns: - The stream ID of the replayed events, or None if no events were found. + The stream ID of the replayed events - the same id `store_event` + received for them - or None if no events were found. The transport + only releases a replay for a stream id it minted itself. """ pass # pragma: no cover +class _MessageChannel: + """One SSE stream's outbound messages: a request's response stream, or the standalone GET stream. + + Every message is first offered to the `EventStore` (so a client that + drops the connection can resume via `Last-Event-ID`), then forwarded to + the SSE response currently attached to the channel, if any. A store that + fails degrades resumability - the failure is logged and the message still + goes out live - it never takes the stream down; `closed` means only that + the stream's life is over (its request finished, or the session ended). + """ + + def __init__(self, stream_id: StreamId, event_store: EventStore | None) -> None: + self.stream_id = stream_id + self._event_store = event_store + self._writer: MemoryObjectSendStream[EventMessage] | None = None + self._reader: MemoryObjectReceiveStream[EventMessage] | None = None + self._closed = False + # Store-then-forward is one atomic step per channel, so the wire order + # of concurrent writers always matches the order the event store saw + # (what a `Last-Event-ID` resume replays from). + self._write_lock = anyio.Lock() + self.terminal: JSONRPCResponse | JSONRPCError | None = None + """The terminal outcome once the request this channel serves has finished.""" + self.finished = anyio.Event() + """Set once `terminal` is recorded, or the channel is finished/closed without one.""" + + @property + def attached(self) -> bool: + """Whether an SSE response is currently draining this channel.""" + return self._writer is not None + + async def write(self, message: JSONRPCMessage) -> bool: + """Store-then-forward one outbound message. Never raises. + + Returns whether the message reached somewhere the client can still get + it - the event store, or the currently attached response. `False` + means it was dropped: the stream is over, or nothing could hold it. + """ + async with self._write_lock: + if self._closed: + logger.debug("dropped message on closed stream %s", self.stream_id) + return False + # Store the event if we have an event store, + # regardless of whether a client is connected + # messages will be replayed on the re-connect + event_id: EventId | None = None + if self._event_store is not None: + try: + event_id = await self._event_store.store_event(self.stream_id, message) + except Exception: + # A broken store costs resumability for this message, not the + # stream: log (its text never reaches the wire) and still + # deliver live below. + logger.exception("EventStore.store_event failed for stream %s", self.stream_id) + else: + logger.debug(f"Stored {event_id} from {self.stream_id}") + if isinstance(message, JSONRPCResponse | JSONRPCError): + self.terminal = message + self.finished.set() + writer = self._writer + if writer is None: + logger.debug( + f"""Request stream {self.stream_id} is not connected + for message. Still processing message as the client + might reconnect and replay.""" + ) + # Retrievable later only if the store took it. + return event_id is not None + try: + await writer.send(EventMessage(message, event_id)) + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + # The response's reader closed under the send; the response's + # own cleanup detaches this attachment. + return event_id is not None + return True + + def attach(self) -> MemoryObjectReceiveStream[EventMessage] | None: + """Attach a fresh SSE response and return the reader it drains. + + Returns `None` when the stream's life is already over, so no response + can attach to a dead channel. Callers check `attached` first where a + second reader is an error (the standalone GET stream); re-attaching + after a detach is how a `Last-Event-ID` reconnect resumes a live stream. + """ + if self._closed: + return None + assert self._writer is None, "every attach site checks `attached` first" + writer, reader = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) + self._writer, self._reader = writer, reader + return reader + + def detach(self, reader: MemoryObjectReceiveStream[EventMessage] | None = None) -> None: + """Detach the current SSE response so the request can carry on without it. + + `reader` scopes the detach to the attachment that reader came from: a + stale response ending must not knock a newer (resumed) attachment off + the channel. With no `reader`, whatever is attached is detached. + """ + if self._writer is None: + return + if reader is not None and reader is not self._reader: + return + self._writer.close() + self._writer, self._reader = None, None + + def close(self) -> None: + """The stream is over: detach any response and drop every further write. + + Serves both request completion (the terminal frame normally ended the + response already; this covers one whose terminal write never landed, so + the client sees the stream close rather than hang) and session + termination. Frames already buffered still drain to their reader. + """ + self._closed = True + self.finished.set() + self.detach() + + +@dataclass +class _HTTPRequestDispatchContext: + """`DispatchContext` for one JSON-RPC message received over streamable HTTP. + + For a request POST, `channel` is that request's response stream: request + scoped notifications, progress, and server-to-client requests all ride it. + For a notification POST there is no request in flight, so the same + operations ride the connection's standalone stream instead. + """ + + transport: TransportContext + _corr: RequestCorrelator[_HTTPRequestDispatchContext] + _channel: _MessageChannel + _request_id: RequestId | None + message_metadata: ServerMessageMetadata | None = None # TODO(maxisbey): remove for Context rework + """The per-request HTTP `Request` and SSE close callbacks the server lifts onto its request context.""" + _progress_token: RequestId | None = None + _closed: bool = False + cancel_requested: anyio.Event = field(default_factory=anyio.Event) + + @property + def request_id(self) -> RequestId | None: + return self._request_id + + @property + def can_send_request(self) -> bool: + return self.transport.can_send_request and not self._closed + + async def notify(self, method: str, params: Mapping[str, Any] | None, opts: CallOptions | None = None) -> None: + if self._closed: + logger.debug("dropped %s: dispatch context closed", method) + return + await self._channel.write(_notification(method, params)) + + async def send_raw_request( + self, + method: str, + params: Mapping[str, Any] | None, + opts: CallOptions | None = None, + ) -> dict[str, Any]: + if not self.can_send_request: + raise NoBackChannelError(method) + return await _call_over_channel(self._corr, self._channel, method, params, opts) + + async def progress(self, progress: float, total: float | None = None, message: str | None = None) -> None: + if self._progress_token is None: + return + params: dict[str, Any] = {"progressToken": self._progress_token, "progress": progress} + if total is not None: + params["total"] = total + if message is not None: + params["message"] = message + await self.notify("notifications/progress", params) + + def close(self) -> None: + self._closed = True + + +class _StandaloneOutbound: + """The connection's `Outbound`: server-initiated messages on the standalone GET stream.""" + + def __init__(self, corr: RequestCorrelator[_HTTPRequestDispatchContext], channel: _MessageChannel) -> None: + self._corr = corr + self._channel = channel + + async def send_raw_request( + self, + method: str, + params: Mapping[str, Any] | None, + opts: CallOptions | None = None, + ) -> dict[str, Any]: + return await _call_over_channel(self._corr, self._channel, method, params, opts) + + async def notify(self, method: str, params: Mapping[str, Any] | None, opts: CallOptions | None = None) -> None: + await self._channel.write(_notification(method, params)) + + +def _notification(method: str, params: Mapping[str, Any] | None) -> JSONRPCNotification: + # Leave `params` unset when None: with `exclude_unset=True` an explicit + # None would serialize as `"params": null`, which JSON-RPC 2.0 forbids. + if params is not None: + return JSONRPCNotification(jsonrpc="2.0", method=method, params=dict(params)) + return JSONRPCNotification(jsonrpc="2.0", method=method) + + +async def _call_over_channel( + corr: RequestCorrelator[_HTTPRequestDispatchContext], + channel: _MessageChannel, + method: str, + params: Mapping[str, Any] | None, + opts: CallOptions | None, +) -> dict[str, Any]: + """Send a server-to-client request on `channel` and await the client's POSTed response. + + The whole abandon policy (timeout, caller cancel, courtesy + `notifications/cancelled` written to the same channel) is the shared + `RequestCorrelator`'s; only the write side is HTTP-specific. + """ + opts = opts or {} + + async def write_request(message: JSONRPCRequest) -> None: + if not await channel.write(message): + # Neither stored nor delivered live: the request can never reach + # the client, so fail the caller (the correlator surfaces + # CONNECTION_CLOSED) rather than await an answer that cannot come. + raise anyio.ClosedResourceError + + async def send_cancel(request_id: RequestId, reason: str) -> None: + await channel.write(_notification("notifications/cancelled", {"requestId": request_id, "reason": reason})) + + return await corr.call( + method, + params, + opts, + write_request=write_request, + send_cancel=send_cancel, + cancel_on_abandon=opts.get("cancel_on_abandon", True), + ) + + class StreamableHTTPServerTransport: """HTTP server transport with event streaming support for MCP. Handles JSON-RPC messages in HTTP POST requests with SSE streaming. - Supports optional JSON responses and session management. + Supports optional JSON responses and session management. One instance + serves one session (or, in stateless mode, one request); the + `StreamableHTTPSessionManager` creates and routes to them. """ - # Server notification streams for POST requests as well as standalone SSE stream - _read_stream_writer: ContextSendStream[SessionMessage | Exception] | None = None - _read_stream: ContextReceiveStream[SessionMessage | Exception] | None = None - _write_stream: ContextSendStream[SessionMessage] | None = None - _write_stream_reader: ContextReceiveStream[SessionMessage] | None = None _security: TransportSecurityMiddleware def __init__( @@ -167,6 +426,9 @@ def __init__( event_store: EventStore | None = None, security_settings: TransportSecuritySettings | None = None, retry_interval: int | None = None, + *, + app: Server[Any] | None = None, + lifespan_state: Any = None, ) -> None: """Initialize a new StreamableHTTP server transport. @@ -183,6 +445,10 @@ def __init__( retry field. When set, the server will send a retry field in SSE priming events to control client reconnection timing for polling behavior. Only used when event_store is provided. + app: The `Server` whose handlers serve this session's requests. Only + the `StreamableHTTPSessionManager` need supply this. + lifespan_state: The server's already-entered lifespan output, shared + across every session by the manager. Raises: ValueError: If the session ID contains invalid characters. @@ -195,23 +461,63 @@ def __init__( self._event_store = event_store self._security = TransportSecurityMiddleware(security_settings) self._retry_interval = retry_interval - self._request_streams: dict[ - RequestId, - tuple[ - MemoryObjectSendStream[EventMessage], - MemoryObjectReceiveStream[EventMessage], - ], - ] = {} - self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[SSEEvent]] = {} + self._app = app + self._lifespan_state = lifespan_state self._terminated = False # Idle timeout cancel scope; managed by the session manager. self.idle_scope: anyio.CancelScope | None = None + # Correlates server-to-client requests with the responses the client + # POSTs back, and lets `notifications/cancelled` find in-flight handlers. + self._corr: RequestCorrelator[_HTTPRequestDispatchContext] = RequestCorrelator() + # Stream ids handed to the event store are minted here, in this + # session's own namespace: distinct from anything a client can name + # and unshared with other sessions on a common store. Stateless + # transports (one per request) get a fresh scope each. + self._stream_scope = mcp_session_id if mcp_session_id is not None else uuid4().hex + # The standalone GET stream: server-initiated messages related to no request. + self._standalone = _MessageChannel(f"{self._stream_scope}:{GET_STREAM_KEY}", event_store) + # In-flight request streams, keyed by their event-store stream id + # (`close_sse_stream()` and `Last-Event-ID` replay both look up here). + self._streams: dict[StreamId, _MessageChannel] = {} + # While an `initialize` is being served, other requests wait for its + # commit: the handshake orders before every later request on the wire. + self._initializing: anyio.Event | None = None + # Session-scoped task group for request handlers (stateful mode). Handlers + # outlive the HTTP request that started them: a dropped connection does + # not cancel a 2025-era request (the client cancels explicitly). + self._task_group: anyio.abc.TaskGroup | None = None + # Whether session-bound work may still be scheduled. Cleared the moment + # the session starts to end - explicit terminate, idle timeout, or + # manager shutdown - before the task group drains its running handlers. + self._accepting = True + self._closed_event = anyio.Event() + # The stateful session's connection state and handler kernel; stateless + # mode builds a born-ready connection per request instead. + self._connection: Connection | None = None + self._runner: ServerRunner[Any] | None = None + if app is not None and mcp_session_id is not None: + outbound = _StandaloneOutbound(self._corr, self._standalone) + self._connection = Connection.for_loop(outbound, session_id=mcp_session_id) + self._runner = ServerRunner(app, self._connection, lifespan_state) @property def is_terminated(self) -> bool: """Check if this transport has been explicitly terminated.""" return self._terminated + def _owns_stream(self, stream_id: StreamId) -> bool: + """Whether an event-store stream id was minted by this transport (this session).""" + return stream_id == self._standalone.stream_id or stream_id.startswith(f"{self._stream_scope}:request:") + + def _request_stream_id(self, request_id: RequestId) -> StreamId: + """The event-store stream id for one request's response stream. + + Minted in this session's namespace with a `request:` infix, so it is + neither reachable from another session sharing the store nor equal to + the standalone stream's id, whatever the client picks as request id. + """ + return f"{self._stream_scope}:request:{request_id}" + def close_sse_stream(self, request_id: RequestId) -> None: """Close SSE connection for a specific request without terminating the stream. @@ -230,15 +536,9 @@ def close_sse_stream(self, request_id: RequestId) -> None: Requires event_store to be configured for events to be stored during the disconnect. """ - writer = self._sse_stream_writers.pop(request_id, None) - if writer: # pragma: no branch - writer.close() - - # Also close and remove request streams - if request_id in self._request_streams: # pragma: no branch - send_stream, receive_stream = self._request_streams.pop(request_id) - send_stream.close() - receive_stream.close() + channel = self._streams.get(self._request_stream_id(request_id)) + if channel is not None: + channel.detach() def close_standalone_sse_stream(self) -> None: """Close the standalone GET SSE stream, triggering client reconnection. @@ -255,25 +555,59 @@ def close_standalone_sse_stream(self) -> None: Requires event_store to be configured for events to be stored during the disconnect. """ - self.close_sse_stream(GET_STREAM_KEY) + self._standalone.detach() + + async def run(self, *, task_status: anyio.abc.TaskStatus[None] = anyio.TASK_STATUS_IGNORED) -> None: + """Host this session's request-handler tasks until the session is terminated. - def _create_session_message( + Request handlers run here rather than inside the HTTP request that + started them, so a client that drops a connection does not cancel its + request (per the 2025-era transport spec). Returns after `terminate()`; + cancelling it (server shutdown or the idle timeout) cancels the + handlers and tears the connection down. + """ + self._require_app() + connection = self._connection + assert connection is not None, "a session-bound transport always has a connection" + try: + async with anyio.create_task_group() as tg: + self._task_group = tg + task_status.started() + try: + await self._closed_event.wait() + finally: + # However the session ends, stop taking work before the + # task group drains the handlers already running. + self._accepting = False + tg.cancel_scope.cancel() + finally: + self._task_group = None + # By now every request handler has finished and released its own + # channel; end the standalone stream and wake anything awaiting a + # client answer (runs on termination and manager shutdown alike). + self._standalone.close() + self._corr.close() + await aclose_shielded(connection) + + def _build_message_metadata( self, - message: JSONRPCRequest, request: Request, request_id: RequestId, protocol_version: str, - ) -> SessionMessage: - """Create a session message with metadata including close_sse_stream callback. + *, + channel: _MessageChannel | None = None, + ) -> ServerMessageMetadata: + """Build the per-request metadata the handler kernel lifts onto its request context. The close_sse_stream callbacks are only provided when the client supports resumability (protocol version >= 2025-11-25). Old clients can't resume if the stream is closed early because they didn't receive a priming event. - Every request carries `on_request_unanswered`, so a request that settles - without a response is still terminated on this era's wire. + With the request's `channel`, the metadata also carries the hook that + terminates a request settling without a response on this era's wire. """ - end_stream = partial(self._terminate_unanswered_request, message.id) - # Only provide close callbacks when client supports resumability + on_request_unanswered = ( + partial(self._settle_unanswered_request, channel, request_id) if channel is not None else None + ) if self._event_store and is_version_at_least(protocol_version, "2025-11-25"): async def close_stream_callback() -> None: @@ -282,23 +616,23 @@ async def close_stream_callback() -> None: async def close_standalone_stream_callback() -> None: self.close_standalone_sse_stream() - metadata = ServerMessageMetadata( + return ServerMessageMetadata( request_context=request, close_sse_stream=close_stream_callback, close_standalone_sse_stream=close_standalone_stream_callback, - on_request_unanswered=end_stream, + on_request_unanswered=on_request_unanswered, ) - else: - metadata = ServerMessageMetadata(request_context=request, on_request_unanswered=end_stream) + return ServerMessageMetadata(request_context=request, on_request_unanswered=on_request_unanswered) - return SessionMessage(message, metadata=metadata) + def _transport_context(self, request: Request, *, can_send_request: bool) -> TransportContext: + return TransportContext(kind="streamable-http", can_send_request=can_send_request, headers=request.headers) async def _mint_priming_event(self, stream_id: StreamId, protocol_version: str) -> SSEEvent | None: """Store the priming cursor for `stream_id` and return its SSE wire form. Called before the request is dispatched so the priming row precedes - anything `message_router` can store for this stream. Returns `None` - when no event store is configured or the client predates 2025-11-25 + anything the handler can store for this stream. Returns `None` when + no event store is configured or the client predates 2025-11-25 (older clients cannot parse the empty-data event). """ if not self._event_store: @@ -311,31 +645,6 @@ async def _mint_priming_event(self, stream_id: StreamId, protocol_version: str) priming_event["retry"] = self._retry_interval return priming_event - async def _run_sse_writer( - self, - request_id: RequestId, - sse_stream_writer: MemoryObjectSendStream[SSEEvent], - request_stream_reader: MemoryObjectReceiveStream[EventMessage], - priming_event: SSEEvent | None, - ) -> None: - """Forward `_request_streams[request_id]` onto the SSE wire for one POST.""" - try: - async with sse_stream_writer, request_stream_reader: - if priming_event is not None: - await sse_stream_writer.send(priming_event) - async for event_message in request_stream_reader: - await sse_stream_writer.send(self._create_event_data(event_message)) - if isinstance(event_message.message, JSONRPCResponse | JSONRPCError): - break - except anyio.ClosedResourceError: # pragma: lax no cover - logger.debug("SSE stream closed by close_sse_stream()") - except Exception: # pragma: lax no cover - logger.exception("Error in SSE writer") - finally: - logger.debug("Closing SSE writer") - self._sse_stream_writers.pop(request_id, None) - await self._clean_up_memory_streams(request_id) - def _create_error_response( self, error_message: str, @@ -401,36 +710,23 @@ def _create_event_data(self, event_message: EventMessage) -> SSEEvent: return event_data - async def _terminate_unanswered_request(self, request_id: RequestId) -> None: - """Terminate a request that settled without a response (e.g. cancelled). - - The 2025-era wire ends a request's stream only with a response for its - id - and stores that response so a resuming client's replay terminates - too - so this era answers a cancelled request with `REQUEST_CANCELLED` - where the dispatcher itself stays silent (the 2026 transports MUST NOT - answer). It is written through the same ordered channel as the request's - other messages, so it cannot overtake anything already queued for it. - """ - assert self._write_stream is not None # a dispatched request implies connect() ran - error = ErrorData(code=REQUEST_CANCELLED, message="Request cancelled") - await self._write_stream.send(SessionMessage(JSONRPCError(jsonrpc="2.0", id=request_id, error=error))) - - async def _clean_up_memory_streams(self, request_id: RequestId) -> None: - """Clean up memory streams for a given request ID.""" - if request_id in self._request_streams: # pragma: no branch - try: - # Close the request stream - await self._request_streams[request_id][0].aclose() - await self._request_streams[request_id][1].aclose() - except Exception: # pragma: no cover - # During cleanup, we catch all exceptions since streams might be in various states - logger.debug("Error closing memory streams - may already be closed") - finally: - # Remove the request stream from the mapping - self._request_streams.pop(request_id, None) + def _sse_headers(self) -> dict[str, str]: + return { + "Cache-Control": "no-cache, no-transform", + "Connection": "keep-alive", + "Content-Type": CONTENT_TYPE_SSE, + **({MCP_SESSION_ID_HEADER: self.mcp_session_id} if self.mcp_session_id else {}), + } async def handle_request(self, scope: Scope, receive: Receive, send: Send) -> None: - """Application entry point that handles all HTTP requests.""" + """Application entry point that handles all HTTP requests. + + Raises: + RuntimeError: The transport was constructed without a server to + dispatch to (`app`); it is created and driven by + `StreamableHTTPSessionManager`. + """ + self._require_app() request = Request(scope, receive) # Validate request headers for DNS rebinding protection @@ -487,11 +783,16 @@ async def _validate_accept_header(self, request: Request, scope: Scope, send: Se return False return True + def _require_app(self) -> Server[Any]: + if self._app is None: + raise RuntimeError( + "StreamableHTTPServerTransport is not bound to a server; " + "it is created and driven by StreamableHTTPSessionManager" + ) + return self._app + async def _handle_post_request(self, scope: Scope, request: Request, receive: Receive, send: Send) -> None: """Handle POST requests containing JSON-RPC messages.""" - writer = self._read_stream_writer - if writer is None: # pragma: no cover - raise ValueError("No read stream writer available. Ensure connect() is called first.") try: # Validate Accept header if not await self._validate_accept_header(request, scope, send): @@ -554,13 +855,13 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re None, HTTPStatus.ACCEPTED, ) - await response(scope, receive, send) - - # Process the message after sending the response - metadata = ServerMessageMetadata(request_context=request) - session_message = SessionMessage(message, metadata=metadata) - await writer.send(session_message) - + try: + await response(scope, receive, send) + finally: + # A body that arrived in full is delivered even when the + # 202 could not be (the client dropped after sending it): + # a lost ack must not lose an answer or a cancellation. + await self._deliver_client_message(request, message) return # Extract protocol version for priming event decision. @@ -572,102 +873,9 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re else request.headers.get(MCP_PROTOCOL_VERSION_HEADER, DEFAULT_NEGOTIATED_VERSION) ) - request_id = str(message.id) + await self._serve_request(scope, request, receive, send, message, protocol_version) - if self.is_json_response_enabled: - self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - request_stream_reader = self._request_streams[request_id][1] - # Process the message - metadata = ServerMessageMetadata( - request_context=request, - on_request_unanswered=partial(self._terminate_unanswered_request, message.id), - ) - session_message = SessionMessage(message, metadata=metadata) - await writer.send(session_message) - try: - # Process messages from the request-specific stream - # We need to collect all messages until we get a response - response_message = None - - # Use similar approach to SSE writer for consistency - async for event_message in request_stream_reader: # pragma: no branch - # If it's a response, this is what we're waiting for - if isinstance(event_message.message, JSONRPCResponse | JSONRPCError): - response_message = event_message.message - break - # For notifications and requests, keep waiting - else: # pragma: no cover - logger.debug(f"received: {event_message.message.method}") - - # At this point we should have a response - if response_message: - # Create JSON response - response = self._create_json_response(response_message) - await response(scope, receive, send) - else: # pragma: no cover - # This shouldn't happen in normal operation - logger.error("No response message received before stream closed") - response = self._create_error_response( - "Error processing request: No response received", - HTTPStatus.INTERNAL_SERVER_ERROR, - ) - await response(scope, receive, send) - except Exception: # pragma: no cover - logger.exception("Error processing JSON response") - response = self._create_error_response( - "Error processing request", - HTTPStatus.INTERNAL_SERVER_ERROR, - INTERNAL_ERROR, - ) - await response(scope, receive, send) - finally: - await self._clean_up_memory_streams(request_id) - else: - # Mint the priming event before any per-request state exists: - # `EventStore.store_event` is user code and may raise, in which - # case the outer handler returns a 500 with nothing to clean up. - # Still strictly precedes dispatch, so storage order == wire order. - priming_event = await self._mint_priming_event(request_id, protocol_version) - - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) - self._sse_stream_writers[request_id] = sse_stream_writer - self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - request_stream_reader = self._request_streams[request_id][1] - - headers = { - "Cache-Control": "no-cache, no-transform", - "Connection": "keep-alive", - "Content-Type": CONTENT_TYPE_SSE, - **({MCP_SESSION_ID_HEADER: self.mcp_session_id} if self.mcp_session_id else {}), - } - response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=partial( - self._run_sse_writer, request_id, sse_stream_writer, request_stream_reader, priming_event - ), - headers=headers, - ) - - # Start the SSE response (this will send headers immediately) - try: - # First send the response to establish the SSE connection - async with anyio.create_task_group() as tg: - tg.start_soon(response, scope, receive, send) - # Then send the message to be processed by the server - session_message = self._create_session_message(message, request, request_id, protocol_version) - await writer.send(session_message) - except Exception: # pragma: lax no cover - logger.exception("SSE response error") - await sse_stream_writer.aclose() - await self._clean_up_memory_streams(request_id) - finally: - await sse_stream_reader.aclose() - - except Exception as err: + except Exception: logger.exception("Error handling POST request") response = self._create_error_response( "Error handling POST request", @@ -675,9 +883,353 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re INTERNAL_ERROR, ) await response(scope, receive, send) - await writer.send(Exception(err)) return + def _session_runner(self) -> ServerRunner[Any]: + """The stateful session's handler kernel; built with the transport's server binding.""" + assert self._runner is not None + return self._runner + + def _stateless_runner(self, request: Request) -> ServerRunner[Any]: + """A born-ready, no-back-channel kernel for one stateless request. + + The `MCP-Protocol-Version` header (or the spec's default when it is + absent) seeds `ctx.protocol_version`; there is no handshake to negotiate it. + """ + protocol_version = request.headers.get(MCP_PROTOCOL_VERSION_HEADER, DEFAULT_NEGOTIATED_VERSION) + connection = Connection.from_envelope(protocol_version, None, None) + return ServerRunner(self._require_app(), connection, self._lifespan_state) + + async def _deliver_client_message( + self, request: Request, message: JSONRPCNotification | JSONRPCResponse | JSONRPCError + ) -> None: + """Handle a POSTed response (to a server-initiated request) or notification, after the 202.""" + if isinstance(message, JSONRPCResponse): + self._corr.resolve(message.id, message.result) + return + if isinstance(message, JSONRPCError): + self._corr.resolve(message.id, message.error) + return + if message.method == "notifications/cancelled": + self._corr.peer_cancel(cancelled_request_id_from_params(message.params), interrupt=True) + elif message.method == "notifications/progress": + delivery = self._corr.progress_callback(message.params) + if delivery is not None: + fn, progress, total, note = delivery + await self._spawn_or_run(fn, progress, total, note) + if self.mcp_session_id is None: + runner = self._stateless_runner(request) + connection = runner.connection + else: + runner = self._session_runner() + connection = None + dctx = _HTTPRequestDispatchContext( + transport=self._transport_context(request, can_send_request=self._can_send_request), + _corr=self._corr, + _channel=self._standalone, + _request_id=None, + message_metadata=ServerMessageMetadata(request_context=request), + ) + + async def _run_notification() -> None: + # `on_notify` contains handler exceptions itself, so a crashing + # notification handler cannot take the session down. + try: + await runner.on_notify(dctx, message.method, message.params) + finally: + if connection is not None: + await aclose_shielded(connection) + + await self._spawn_or_run(_run_notification) + + @property + def _can_send_request(self) -> bool: + """Whether a request handler on this transport has a back-channel for server-to-client requests. + + JSON-response mode has none: the request's response is one JSON body, + so a nested elicitation or sampling request has no stream to ride. + Stateless mode additionally lacks a session, so no POST of the client's + answer could be correlated back to a waiting handler. + """ + return self.mcp_session_id is not None and not self.is_json_response_enabled + + async def _spawn_or_run(self, fn: Callable[..., Awaitable[None]], *args: Any) -> None: + """Run `fn(*args)`: on the session task group when session-bound, inline when stateless. + + A session-bound transport whose session has ended drops the work: the + session is being torn down, so no handler may run against it. + """ + if self.mcp_session_id is None: + await fn(*args) + return + tg = self._task_group + if tg is None or not self._accepting: + logger.debug("dropped work for ended session %s", self.mcp_session_id) + return + tg.start_soon(fn, *args) + + async def _serve_request( + self, + scope: Scope, + request: Request, + receive: Receive, + send: Send, + message: JSONRPCRequest, + protocol_version: str, + ) -> None: + """Dispatch one POSTed JSON-RPC request and stream its response.""" + request_id = message.id + stream_id = self._request_stream_id(request_id) + stateful = self.mcp_session_id is not None + + # A request other than `initialize` waits for a handshake in progress + # to commit: over one stream the read loop parked to guarantee this; + # over HTTP the requests are concurrent, so the transport keeps the + # same order explicitly. This handshake's gate exists before its + # first await. + initialize_gate: anyio.Event | None = None + if stateful and message.method == "initialize": + initialize_gate = anyio.Event() + self._initializing = initialize_gate + elif stateful and (gate := self._initializing) is not None: + await gate.wait() + + # From here the gate must be released whatever becomes of this + # request: by the handler task once it exists, else by this frame. + handler_started = False + try: + handler_started = await self._start_request( + scope, request, receive, send, message, protocol_version, stream_id, initialize_gate + ) + finally: + if initialize_gate is not None and not handler_started: + self._release_initialize_gate(initialize_gate) + + def _release_initialize_gate(self, gate: anyio.Event) -> None: + """Let the requests held behind this handshake proceed, and clear it if still current.""" + gate.set() + if self._initializing is gate: + self._initializing = None + + async def _start_request( + self, + scope: Scope, + request: Request, + receive: Receive, + send: Send, + message: JSONRPCRequest, + protocol_version: str, + stream_id: StreamId, + initialize_gate: anyio.Event | None, + ) -> bool: + """Register and start one request; returns whether a handler task took over its lifecycle.""" + request_id = message.id + stateful = self.mcp_session_id is not None + + # Mint the priming event before any per-request state exists: + # `EventStore.store_event` is user code and may raise, in which + # case the outer handler returns a 500 with nothing to clean up. + # Still strictly precedes dispatch, so storage order == wire order. + priming_event = ( + None if self.is_json_response_enabled else await self._mint_priming_event(stream_id, protocol_version) + ) + + # The session may have ended (DELETE, idle timeout, manager shutdown) + # while this request was suspended above; a session-bound transport must + # refuse the request instead of running it against a dead session. No + # await from here to the dispatch, so the answer holds when we act on it. + session_task_group = self._task_group + if stateful and (not self._accepting or session_task_group is None): + response = self._create_error_response( + "Not Found: Session has been terminated", + HTTPStatus.NOT_FOUND, + ) + await response(scope, receive, send) + return False + + channel = _MessageChannel(stream_id, self._event_store) + self._streams[stream_id] = channel + # Attach the response's writer before the handler starts, so nothing + # the handler emits early lands on an unattached channel. + reader = None if self.is_json_response_enabled else channel.attach() + + if stateful: + runner = self._session_runner() + connection = None + else: + runner = self._stateless_runner(request) + connection = runner.connection + dctx = _HTTPRequestDispatchContext( + transport=self._transport_context(request, can_send_request=self._can_send_request), + _corr=self._corr, + _channel=channel, + _request_id=request_id, + message_metadata=self._build_message_metadata(request, request_id, protocol_version, channel=channel), + _progress_token=progress_token_from_params(message.params), + ) + cancel_scope = anyio.CancelScope() + self._corr.enter_inbound(request_id, cancel_scope, dctx) + + async def _run_handler() -> None: + try: + # `serve_inbound` contains handler exceptions and `channel.write` + # never raises, so this task always completes on its own. + await self._corr.serve_inbound( + request_id, + dctx, + cancel_scope, + partial(runner.on_request, dctx, message.method, message.params), + write_result=partial(self._write_result, channel, request_id), + write_error=partial(self._write_error, channel, request_id), + settle_unanswered=dctx.message_metadata.on_request_unanswered if dctx.message_metadata else None, + ) + finally: + if initialize_gate is not None: + self._release_initialize_gate(initialize_gate) + # The channel stays registered until the handler is done, so a + # `Last-Event-ID` reconnect can re-attach while it still runs. + channel.close() + if self._streams.get(stream_id) is channel: + del self._streams[stream_id] + if connection is not None: + await aclose_shielded(connection) + + if session_task_group is not None: + # Session-scoped: the handler outlives this HTTP request. A client + # that drops the connection is not cancelling the request (it may + # resume via Last-Event-ID); it cancels by POSTing notifications/cancelled. + session_task_group.start_soon(_run_handler) + if reader is None: + await self._respond_json(scope, receive, send, channel) + else: + await self._respond_sse(scope, receive, send, channel, reader, priming_event) + else: + # Stateless: this request is the whole connection, so the handler's + # lifetime is the response's - it is cancelled once the response + # ends (result delivered, or the client went away). + async with anyio.create_task_group() as tg: + tg.start_soon(_run_handler) + if reader is None: + await self._respond_json(scope, receive, send, channel) + else: + await self._respond_sse(scope, receive, send, channel, reader, priming_event) + tg.cancel_scope.cancel() + return True + + async def _settle_unanswered_request(self, channel: _MessageChannel, request_id: RequestId) -> None: + """Terminate a request that settled without a response (e.g. it was cancelled). + + The 2025-era wire ends a request's stream only with a response for its + id - and stores that response so a resuming client's replay terminates + too - so this era answers a cancelled request with `REQUEST_CANCELLED` + where the dispatch layer itself stays silent (the 2026 transports MUST + NOT answer). It goes through the request's own ordered channel, so it + cannot overtake anything already queued for it. + """ + await self._write_error(channel, request_id, ErrorData(code=REQUEST_CANCELLED, message="Request cancelled")) + + async def _write_result(self, channel: _MessageChannel, request_id: RequestId, result: dict[str, Any]) -> None: + await channel.write(JSONRPCResponse(jsonrpc="2.0", id=request_id, result=result)) + + async def _write_error(self, channel: _MessageChannel, request_id: RequestId, error: ErrorData) -> None: + await channel.write(JSONRPCError(jsonrpc="2.0", id=request_id, error=error)) + + async def _respond_json(self, scope: Scope, receive: Receive, send: Send, channel: _MessageChannel) -> None: + """Wait for the request's terminal message and send it as one JSON body.""" + await channel.finished.wait() + response_message = channel.terminal + if response_message is not None: + response = self._create_json_response(response_message) + elif self._terminated: + # The session ended underneath the request; it gets the same + # answer every request to a terminated session gets. + response = self._create_error_response( + "Not Found: Session has been terminated", + HTTPStatus.NOT_FOUND, + ) + else: # pragma: lax no cover + # The request finished without recording an answer (a wedged store + # can outlast the shutdown write bound). Nothing to send but a 500. + logger.error("No response message received before stream closed") + response = self._create_error_response( + "Error processing request: No response received", + HTTPStatus.INTERNAL_SERVER_ERROR, + ) + await response(scope, receive, send) + + async def _respond_sse( + self, + scope: Scope, + receive: Receive, + send: Send, + channel: _MessageChannel, + reader: MemoryObjectReceiveStream[EventMessage] | None, + priming_event: SSEEvent | None, + ) -> None: + """Stream the request's channel as this POST's SSE response, until the response frame passes.""" + assert reader is not None, "a freshly created request channel always attaches" + await self._run_sse_response( + scope, receive, send, partial(self._pump_channel, channel, reader, priming_event, stop_at_response=True) + ) + # The client is gone (disconnect or delivered response): detach so the + # handler carries on writing to the store alone. + channel.detach(reader) + + async def _run_sse_response( + self, + scope: Scope, + receive: Receive, + send: Send, + data_sender: Callable[[MemoryObjectSendStream[SSEEvent]], Coroutine[Any, Any, None]], + ) -> None: + """Run one SSE response fed by `data_sender`, the single containment site for all of them. + + `data_sender(sse_send)` writes the events (a channel pump, or a replay + followed by a pump). An error escaping the started response is logged + here and goes no further: the response already began, so nothing may + answer this request a second time. + """ + sse_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) + response = EventSourceResponse( + content=sse_recv, + data_sender_callable=partial(data_sender, sse_send), + headers=self._sse_headers(), + ) + try: + await response(scope, receive, send) + except Exception: # pragma: lax no cover + logger.exception("Error in SSE response") + finally: + await sse_send.aclose() + await sse_recv.aclose() + + async def _pump_channel( + self, + channel: _MessageChannel, + reader: MemoryObjectReceiveStream[EventMessage], + priming_event: SSEEvent | None, + sse_send: MemoryObjectSendStream[SSEEvent], + *, + stop_at_response: bool, + ) -> None: + """Forward one attachment of `channel` onto an SSE response's event queue. + + Runs as sse-starlette's data sender, so a client disconnect cancels it + along with the response; the `finally` detaches this attachment (never + a newer one that a `Last-Event-ID` reconnect may have installed). + """ + try: + async with sse_send, reader: + if priming_event is not None: + await sse_send.send(priming_event) + async for event_message in reader: + await sse_send.send(self._create_event_data(event_message)) + if stop_at_response and isinstance(event_message.message, JSONRPCResponse | JSONRPCError): + break + finally: + logger.debug("Closing SSE writer") + channel.detach(reader) + async def _handle_get_request(self, request: Request, send: Send) -> None: """Handle GET request to establish SSE. @@ -685,10 +1237,6 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: first sending data via HTTP POST. The server can send JSON-RPC requests and notifications on this stream. """ - writer = self._read_stream_writer - if writer is None: # pragma: no cover - raise ValueError("No read stream writer available. Ensure connect() is called first.") - # Validate Accept header - must include text/event-stream _, has_sse = check_accept_headers(request) @@ -704,21 +1252,12 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: return # Handle resumability: check for Last-Event-ID header - if last_event_id := request.headers.get(LAST_EVENT_ID_HEADER): + if self._event_store and (last_event_id := request.headers.get(LAST_EVENT_ID_HEADER)): await self._replay_events(last_event_id, request, send) return - headers = { - "Cache-Control": "no-cache, no-transform", - "Connection": "keep-alive", - "Content-Type": CONTENT_TYPE_SSE, - } - - if self.mcp_session_id: # pragma: no branch - headers[MCP_SESSION_ID_HEADER] = self.mcp_session_id - # Check if we already have an active GET stream - if GET_STREAM_KEY in self._request_streams: + if self._standalone.attached: response = self._create_error_response( "Conflict: Only one SSE stream is allowed per session", HTTPStatus.CONFLICT, @@ -726,54 +1265,25 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: await response(request.scope, request.receive, send) return - # Create SSE stream - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) - - async def standalone_sse_writer(): - try: - # Create a standalone message stream for server-initiated messages - - self._request_streams[GET_STREAM_KEY] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - standalone_stream_reader = self._request_streams[GET_STREAM_KEY][1] - - async with sse_stream_writer, standalone_stream_reader: - # Process messages from the standalone stream - async for event_message in standalone_stream_reader: - # For the standalone stream, we handle: - # - JSONRPCNotification (server sends notifications to client) - # - JSONRPCRequest (server sends requests to client) - # We should NOT receive JSONRPCResponse - - # Send the message via SSE - event_data = self._create_event_data(event_message) - await sse_stream_writer.send(event_data) - except anyio.ClosedResourceError: - # Session teardown can close the stream while the writer is between dequeues. - pass - except Exception: - logger.exception("Error in standalone SSE writer") # pragma: no cover - finally: - logger.debug("Closing standalone SSE writer") - await self._clean_up_memory_streams(GET_STREAM_KEY) - - # Create and start EventSourceResponse - response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=standalone_sse_writer, - headers=headers, - ) - + reader = self._standalone.attach() + if reader is None: # pragma: lax no cover + # The session was terminated between the entry check and here. + response = self._create_error_response( + "Not Found: Session has been terminated", + HTTPStatus.NOT_FOUND, + ) + await response(request.scope, request.receive, send) + return try: # This will send headers immediately and establish the SSE connection - await response(request.scope, request.receive, send) - except Exception: # pragma: lax no cover - logger.exception("Error in standalone SSE response") - await self._clean_up_memory_streams(GET_STREAM_KEY) + await self._run_sse_response( + request.scope, + request.receive, + send, + partial(self._pump_channel, self._standalone, reader, None, stop_at_response=False), + ) finally: - await sse_stream_writer.aclose() - await sse_stream_reader.aclose() + self._standalone.detach(reader) async def _handle_delete_request(self, request: Request, send: Send) -> None: """Handle DELETE requests for explicit session termination.""" @@ -805,29 +1315,20 @@ async def terminate(self) -> None: """ self._terminated = True + self._accepting = False logger.info(f"Terminating session: {self.mcp_session_id}") - # We need a copy of the keys to avoid modification during iteration - request_stream_keys = list(self._request_streams.keys()) - - # Close all request streams asynchronously - for key in request_stream_keys: - await self._clean_up_memory_streams(key) - - # Clear the request streams dictionary immediately - self._request_streams.clear() - try: - if self._read_stream_writer is not None: # pragma: no branch - await self._read_stream_writer.aclose() - if self._read_stream is not None: # pragma: no branch - await self._read_stream.aclose() - if self._write_stream_reader is not None: # pragma: no branch - await self._write_stream_reader.aclose() - if self._write_stream is not None: # pragma: no branch - await self._write_stream.aclose() - except Exception as e: # pragma: no cover - # During cleanup, we catch all exceptions since streams might be in various states - logger.debug(f"Error closing streams: {e}") + # Close every open response stream, wake anything awaiting a + # client answer, and cancel in-flight handlers. + for channel in list(self._streams.values()): + channel.close() + self._streams.clear() + self._standalone.close() + self._corr.close() + self._corr.cancel_all_inbound() + # Release the session task, which cancels any handler still running + # and closes the connection's exit stack. + self._closed_event.set() async def _handle_unsupported_request(self, request: Request, send: Send) -> None: """Handle unsupported HTTP methods.""" @@ -890,80 +1391,78 @@ async def _replay_events(self, last_event_id: str, request: Request, send: Send) return # pragma: no cover try: - headers = { - "Cache-Control": "no-cache, no-transform", - "Connection": "keep-alive", - "Content-Type": CONTENT_TYPE_SSE, - } - - if self.mcp_session_id: # pragma: no branch - headers[MCP_SESSION_ID_HEADER] = self.mcp_session_id - # The manager only routes supported (or absent) header values to this transport replay_protocol_version = request.headers.get(MCP_PROTOCOL_VERSION_HEADER, DEFAULT_NEGOTIATED_VERSION) - # Create SSE stream for replay - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) - - async def replay_sender(): + async def replay_then_tail(sse_send: MemoryObjectSendStream[SSEEvent]) -> None: try: - async with sse_stream_writer: - # Define an async callback for sending events - async def send_event(event_message: EventMessage) -> None: - event_data = self._create_event_data(event_message) - await sse_stream_writer.send(event_data) + async with sse_send: + # Buffer the replay until the store names its stream: the + # event id came from the client, so only a stream in this + # session's namespace is allowed onto the wire. + replayed: list[EventMessage] = [] + + async def collect_event(event_message: EventMessage) -> None: + replayed.append(event_message) # Replay past events and get the stream ID - stream_id = await event_store.replay_events_after(last_event_id, send_event) - - # If stream ID not in mapping, create it - if stream_id and stream_id not in self._request_streams: # pragma: no branch - try: - # Register SSE writer so close_sse_stream() can close it - self._sse_stream_writers[stream_id] = sse_stream_writer - - # Prime the resumed connection so the client sees the stream - # is re-registered. The replay→live-tail ordering window here - # is pre-existing and tracked separately. - priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) - if priming_event is not None: - await sse_stream_writer.send(priming_event) - - # Create new request streams for this connection - self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - msg_reader = self._request_streams[stream_id][1] - - # Forward messages to SSE - async with msg_reader: - async for event_message in msg_reader: - event_data = self._create_event_data(event_message) - - await sse_stream_writer.send(event_data) - finally: - self._sse_stream_writers.pop(stream_id, None) - await self._clean_up_memory_streams(stream_id) - except anyio.ClosedResourceError: # pragma: lax no cover - # Expected when close_sse_stream() is called - logger.debug("Replay SSE stream closed by close_sse_stream()") - except Exception: # pragma: lax no cover + stream_id = await event_store.replay_events_after(last_event_id, collect_event) + if not stream_id: + return + if not self._owns_stream(stream_id): + logger.warning( + "Refusing to replay foreign stream %r on session %s", stream_id, self.mcp_session_id + ) + return + for event_message in replayed: + await sse_send.send(self._create_event_data(event_message)) + + # Live-tail the stream if it is still open and no response + # is currently attached to it: the `close_sse_stream()` + # polling reconnect, and a client resuming a dropped connection. + if stream_id == self._standalone.stream_id: + channel = self._standalone + else: + channel = self._streams.get(stream_id) + if channel is None or channel.attached: + return + + # Attach first, so anything the still-running request emits + # from here on is buffered for this response rather than + # only stored. The replay→live-tail ordering window (frames + # stored between the replay read and the attach) is pre-existing + # and tracked separately. + reader = channel.attach() + if reader is None: + # The stream ended (session terminated) while the store + # was read; there is nothing left to tail. + return + try: + # Prime the resumed connection so the client sees the + # stream is re-registered. + priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) + + # Forward messages to SSE: a request's stream ends after + # its response frame; the standalone stream carries no + # response and tails until the client leaves again. + await self._pump_channel( + channel, + reader, + priming_event, + sse_send, + stop_at_response=channel is not self._standalone, + ) + finally: + channel.detach(reader) + # The pump closes the reader it drained; this covers a + # priming failure that never handed the reader over. + reader.close() + except Exception: + # `replay_events_after` is user code; a failing replay ends this response only. logger.exception("Error in replay sender") # Create and start EventSourceResponse - response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=replay_sender, - headers=headers, - ) - - try: - await response(request.scope, request.receive, send) - except Exception: # pragma: lax no cover - logger.exception("Error in replay response") - finally: - await sse_stream_writer.aclose() - await sse_stream_reader.aclose() + await self._run_sse_response(request.scope, request.receive, send, replay_then_tail) except Exception: # pragma: lax no cover logger.exception("Error replaying events") @@ -973,107 +1472,3 @@ async def send_event(event_message: EventMessage) -> None: INTERNAL_ERROR, ) await response(request.scope, request.receive, send) - - @asynccontextmanager - async def connect( - self, - ) -> AsyncGenerator[ - tuple[ - ReadStream[SessionMessage | Exception], - WriteStream[SessionMessage], - ], - None, - ]: - """Context manager that provides read and write streams for a connection. - - Yields: - Tuple of (read_stream, write_stream) for bidirectional communication - """ - - # Create the memory streams for this connection - - read_stream_writer, read_stream = create_context_streams[SessionMessage | Exception](0) - write_stream, write_stream_reader = create_context_streams[SessionMessage](0) - - # Store the streams - self._read_stream_writer = read_stream_writer - self._read_stream = read_stream - self._write_stream_reader = write_stream_reader - self._write_stream = write_stream - - # Start a task group for message routing - async with anyio.create_task_group() as tg: - # Create a message router that distributes messages to request streams - async def message_router(): - try: - async for session_message in write_stream_reader: # pragma: no branch - # Determine which request stream(s) should receive this message - message = session_message.message - target_request_id = None - # Check if this is a response with a known request id. - # Null-id errors (e.g., parse errors) fall through to - # the GET stream since they can't be correlated. - if isinstance(message, JSONRPCResponse | JSONRPCError) and message.id is not None: - target_request_id = str(message.id) - # Extract related_request_id from meta if it exists - elif ( - session_message.metadata is not None - and isinstance( - session_message.metadata, - ServerMessageMetadata, - ) - and session_message.metadata.related_request_id is not None - ): - target_request_id = str(session_message.metadata.related_request_id) - - request_stream_id = target_request_id if target_request_id is not None else GET_STREAM_KEY - - # Store the event if we have an event store, - # regardless of whether a client is connected - # messages will be replayed on the re-connect - event_id = None - if self._event_store: - event_id = await self._event_store.store_event(request_stream_id, message) - logger.debug(f"Stored {event_id} from {request_stream_id}") - - if request_stream_id in self._request_streams: - try: - # Send both the message and the event ID - await self._request_streams[request_stream_id][0].send(EventMessage(message, event_id)) - except (anyio.BrokenResourceError, anyio.ClosedResourceError): # pragma: no cover - # Stream might be closed, remove from registry - self._request_streams.pop(request_stream_id, None) - else: - logger.debug( - f"""Request stream {request_stream_id} not found - for message. Still processing message as the client - might reconnect and replay.""" - ) - except anyio.ClosedResourceError: - if self._terminated: # pragma: lax no cover - logger.debug("Read stream closed by client") - else: - logger.exception("Unexpected closure of read stream in message router") - except Exception: # pragma: lax no cover - logger.exception("Error in message router") - - # Start the message router - tg.start_soon(message_router) - - try: - # Yield the streams for the caller to use - yield read_stream, write_stream - finally: - for stream_id in list(self._request_streams.keys()): - await self._clean_up_memory_streams(stream_id) - self._request_streams.clear() - - # Clean up the read and write streams - try: - await read_stream_writer.aclose() - await read_stream.aclose() - await write_stream_reader.aclose() - await write_stream.aclose() - except Exception as e: # pragma: no cover - # During cleanup, we catch all exceptions since streams might be in various states - logger.debug(f"Error closing streams: {e}") diff --git a/src/mcp/server/streamable_http_manager.py b/src/mcp/server/streamable_http_manager.py index 31f587ee66..3316856d21 100644 --- a/src/mcp/server/streamable_http_manager.py +++ b/src/mcp/server/streamable_http_manager.py @@ -11,7 +11,7 @@ import anyio from anyio.abc import TaskStatus -from mcp_types import DEFAULT_NEGOTIATED_VERSION, INVALID_REQUEST, ErrorData, JSONRPCError +from mcp_types import INVALID_REQUEST, ErrorData, JSONRPCError from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS from starlette.datastructures import Headers from starlette.requests import Request @@ -20,14 +20,10 @@ from mcp.server._streamable_http_modern import handle_modern_request from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context -from mcp.server.connection import Connection -from mcp.server.runner import serve_connection, serve_loop from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, EventStore, StreamableHTTPServerTransport from mcp.server.transport_security import TransportSecuritySettings from mcp.shared._compat import resync_tracer from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER -from mcp.shared.jsonrpc_dispatcher import JSONRPCDispatcher -from mcp.shared.transport_context import TransportContext if TYPE_CHECKING: from mcp.server.lowlevel.server import Server @@ -188,60 +184,25 @@ async def _handle_request(self, scope: Scope, receive: Receive, send: Send) -> N # Dispatch to the appropriate handler if self.stateless: - await self._handle_stateless_request(pv, scope, receive, send) + await self._handle_stateless_request(scope, receive, send) else: await self._handle_stateful_request(scope, receive, send) - async def _handle_stateless_request( - self, protocol_version_hint: str | None, scope: Scope, receive: Receive, send: Send - ) -> None: + async def _handle_stateless_request(self, scope: Scope, receive: Receive, send: Send) -> None: """Process request in stateless mode - creating a new transport for each request.""" logger.debug("Stateless mode: Creating new transport for this request") - # No session ID needed in stateless mode + # No session ID needed in stateless mode: the transport serves this one + # request with a born-ready connection (no `initialize`, no standalone + # GET stream) and no back-channel for server-to-client requests. http_transport = StreamableHTTPServerTransport( mcp_session_id=None, # No session tracking in stateless mode is_json_response_enabled=self.json_response, event_store=None, # No event store in stateless mode security_settings=self.security_settings, + app=self.app, + lifespan_state=self._lifespan_state, ) - # Start server in a new task - async def run_stateless_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORED): - async with http_transport.connect() as streams: - read_stream, write_stream = streams - task_status.started() - dispatcher: JSONRPCDispatcher[TransportContext] = JSONRPCDispatcher( - read_stream, - write_stream, - inline_methods=frozenset({"initialize"}), - # No session ID means a server-to-client request can be - # written to this POST's response stream, but the client's - # reply has nowhere to land — `can_send_request=False` - # makes the per-request channel raise `NoBackChannelError` - # for requests while still allowing notifications. - transport_builder=lambda _md: TransportContext(kind="streamable-http", can_send_request=False), - ) - # Born-ready, no standalone channel: the legacy stateless path - # never opens a GET stream and need not see `initialize`. The - # header (or the spec's default-absent value) seeds - # `ctx.protocol_version`. - connection = Connection.from_envelope( - protocol_version_hint if protocol_version_hint is not None else DEFAULT_NEGOTIATED_VERSION, - None, - None, - ) - try: - await serve_connection( - self.app, dispatcher, connection=connection, lifespan_state=self._lifespan_state - ) - except Exception: # pragma: lax no cover - logger.exception("Stateless session crashed") - - # Assert task group is not None for type checking - assert self._task_group is not None - # Start the server task - await self._task_group.start(run_stateless_server) - # Handle the HTTP request and return the response await http_transport.handle_request(scope, receive, send) @@ -294,6 +255,10 @@ async def _handle_stateful_request(self, scope: Scope, receive: Receive, send: S event_store=self.event_store, # May be None (no resumability) security_settings=self.security_settings, retry_interval=self.retry_interval, + # The manager owns the lifespan (entered once in `run()`), + # so the transport serves every request off that state. + app=self.app, + lifespan_state=self._lifespan_state, ) assert http_transport.mcp_session_id is not None @@ -302,53 +267,41 @@ async def _handle_stateful_request(self, scope: Scope, receive: Receive, send: S self._server_instances[http_transport.mcp_session_id] = http_transport logger.info(f"Created new transport with session ID: {new_session_id}") - # Define the server runner + # Define the session task: hosts the session's request handlers, + # which outlive the HTTP requests that started them. async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORED) -> None: - async with http_transport.connect() as streams: - read_stream, write_stream = streams - task_status.started() - try: - # Use a cancel scope for idle timeout — when the - # deadline passes the scope cancels the loop and - # execution continues after the ``with`` block. - # Incoming requests push the deadline forward. - idle_scope = anyio.CancelScope() - if self.session_idle_timeout is not None: - idle_scope.deadline = anyio.current_time() + self.session_idle_timeout - http_transport.idle_scope = idle_scope - - with idle_scope: - # Drive via `serve_loop` (not `Server.run()`) so the - # manager's already-entered lifespan is reused - # rather than re-entered per session. - await serve_loop( - self.app, - read_stream, - write_stream, - lifespan_state=self._lifespan_state, - session_id=http_transport.mcp_session_id, - ) - - if idle_scope.cancelled_caught: - assert http_transport.mcp_session_id is not None - logger.info(f"Session {http_transport.mcp_session_id} idle timeout") - self._server_instances.pop(http_transport.mcp_session_id, None) - self._session_owners.pop(http_transport.mcp_session_id, None) - await http_transport.terminate() - except Exception: - logger.exception(f"Session {http_transport.mcp_session_id} crashed") - finally: - if ( # pragma: no branch - http_transport.mcp_session_id - and http_transport.mcp_session_id in self._server_instances - and not http_transport.is_terminated - ): - logger.info( - "Cleaning up crashed session " - f"{http_transport.mcp_session_id} from active instances." - ) - del self._server_instances[http_transport.mcp_session_id] - self._session_owners.pop(http_transport.mcp_session_id, None) + try: + # Use a cancel scope for idle timeout — when the + # deadline passes the scope cancels the session and + # execution continues after the ``with`` block. + # Incoming requests push the deadline forward. + idle_scope = anyio.CancelScope() + if self.session_idle_timeout is not None: + idle_scope.deadline = anyio.current_time() + self.session_idle_timeout + http_transport.idle_scope = idle_scope + + with idle_scope: + await http_transport.run(task_status=task_status) + + if idle_scope.cancelled_caught: + assert http_transport.mcp_session_id is not None + logger.info(f"Session {http_transport.mcp_session_id} idle timeout") + self._server_instances.pop(http_transport.mcp_session_id, None) + self._session_owners.pop(http_transport.mcp_session_id, None) + await http_transport.terminate() + except Exception: + logger.exception(f"Session {http_transport.mcp_session_id} crashed") + finally: + if ( # pragma: no branch + http_transport.mcp_session_id + and http_transport.mcp_session_id in self._server_instances + and not http_transport.is_terminated + ): + logger.info( + f"Cleaning up crashed session {http_transport.mcp_session_id} from active instances." + ) + del self._server_instances[http_transport.mcp_session_id] + self._session_owners.pop(http_transport.mcp_session_id, None) # Assert task group is not None for type checking assert self._task_group is not None diff --git a/src/mcp/shared/_correlation.py b/src/mcp/shared/_correlation.py new file mode 100644 index 0000000000..06df393990 --- /dev/null +++ b/src/mcp/shared/_correlation.py @@ -0,0 +1,491 @@ +"""Request correlation kernel shared by every JSON-RPC-shaped peer. + +`RequestCorrelator` owns the two tables a peer needs regardless of how +messages travel: outbound requests awaiting the peer's response +(`pending`), and inbound requests currently being handled (`in_flight`), +together with everything that hangs off them - request-id minting and the +collision domain, progress routing, peer cancellation, the courtesy-cancel +policy on abandon, the connection-closed fan-out, and the single +exception-to-wire boundary for inbound handlers. + +It knows nothing about framing or streams. Callers supply the write side +as callables: `JSONRPCDispatcher` writes `SessionMessage`s onto its stream +pair; the streamable-HTTP transport writes onto a request's response +channel. That is the whole difference between those transports at this +layer, so the semantics live here exactly once. +""" + +from __future__ import annotations + +import logging +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from functools import partial +from typing import Any, Generic, Protocol + +import anyio +import anyio.lowlevel +from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream +from mcp_types import ( + CONNECTION_CLOSED, + INVALID_PARAMS, + REQUEST_TIMEOUT, + ErrorData, + JSONRPCRequest, + RequestId, +) +from opentelemetry.trace import SpanKind +from pydantic import ValidationError +from typing_extensions import TypeVar + +from mcp.shared._otel import inject_trace_context, otel_span +from mcp.shared.dispatcher import CallOptions, ProgressFnT, coerce_request_id +from mcp.shared.exceptions import MCPError + +__all__ = [ + "InFlight", + "Outcome", + "Pending", + "RequestCorrelator", + "handler_exception_to_error_data", +] + +logger = logging.getLogger(__name__) + +_ABANDON_WRITE_TIMEOUT: float = 5 +"""Bound for courtesy-cancel writes on the abandon paths; the caller-cancel +arm shields its write, so a wedged transport would otherwise hang it uncancellably.""" + +_SHUTDOWN_WRITE_TIMEOUT: float = 1 +"""Tighter bound for the shutdown-arm error write so a wedged transport can't hold session close.""" + +Outcome = dict[str, Any] | ErrorData +"""A request's terminal outcome: the result dict, or the peer's error.""" + + +def handler_exception_to_error_data(exc: BaseException) -> ErrorData | None: + """Map a handler-raised exception to its wire `ErrorData`. + + The two rungs every peer shares: an `MCPError` carries its own + `ErrorData`; a pydantic `ValidationError` is the spec's INVALID_PARAMS + with empty ``data`` (no pydantic text on the wire). Returns ``None`` for + any other exception so each caller applies its own catch-all - + `serve_inbound` currently pins ``code=0`` for v1 compat, + the modern HTTP entry uses `INTERNAL_ERROR`. + """ + if isinstance(exc, MCPError): + return exc.error + if isinstance(exc, ValidationError): + return ErrorData(code=INVALID_PARAMS, message="Invalid request parameters", data="") + return None + + +class _CancelObserver(Protocol): + """The slice of a `DispatchContext` the in-flight table drives on peer cancel.""" + + cancel_requested: anyio.Event + + def close(self) -> None: ... + + +DctxT = TypeVar("DctxT", bound=_CancelObserver, default=_CancelObserver) + + +@dataclass(slots=True) +class Pending: + """An outbound request awaiting its response.""" + + send: MemoryObjectSendStream[Outcome] + receive: MemoryObjectReceiveStream[Outcome] + on_progress: ProgressFnT | None = None + + +@dataclass(slots=True) +class InFlight(Generic[DctxT]): + """An inbound request currently being handled.""" + + scope: anyio.CancelScope + dctx: DctxT + + +def _shielded_progress(fn: ProgressFnT) -> ProgressFnT: + """Wrap a user progress callback so an exception can't cancel the caller's task group.""" + + async def _wrapped(progress: float, total: float | None, message: str | None) -> None: + try: + await fn(progress, total, message) + except Exception: + logger.exception("progress callback raised") + + return _wrapped + + +async def final_write( + write: Callable[[], Awaitable[None]], + *, + shield: bool, + timeout: float, + describe: str, +) -> None: + """Attempt one last write under the shared abandon/teardown policy. + + `shield=True` is for arms already inside a cancelled scope (a bare + `await` would re-raise); the bound keeps a wedged transport write + from becoming an uncancellable hang. + """ + with anyio.move_on_after(timeout, shield=shield) as scope: + await write() + if scope.cancelled_caught: + logger.warning("%s gave up: transport write blocked", describe) + + +class RequestCorrelator(Generic[DctxT]): + """Request-id correlation for one peer connection, in both directions. + + Outbound (`call`, `resolve`, `progress_callback`): requests this side + sent that await the peer's response, keyed by the coerced request id + (the collision domain `coerce_request_id` defines - `"7"` and `7` are + one id even where the wire carries the value verbatim). + + Inbound (`enter_inbound`, `serve_inbound`, `peer_cancel`, + `cancel_all_inbound`): requests the peer sent that this side is + handling, so a `notifications/cancelled` from the peer (or a local + shutdown) can interrupt exactly the right handler. + + `close()` is single-shot: once closed, `call` raises `MCPError` + (`CONNECTION_CLOSED`) and every parked waiter is woken with the same. + """ + + def __init__(self) -> None: + self.pending: dict[RequestId, Pending] = {} + """Outbound requests awaiting a response, keyed by coerced request id.""" + self.in_flight: dict[RequestId, InFlight[DctxT]] = {} + """Inbound requests being handled, keyed by coerced request id.""" + self._next_id = 0 + self._closed = False + + @property + def closed(self) -> bool: + """True once `close()` has run; `call` refuses and waiters were woken.""" + return self._closed + + def close(self) -> None: + """Mark closed and wake every outbound waiter with `CONNECTION_CLOSED`. Idempotent, synchronous.""" + self._closed = True + self.fan_out_closed() + + # ------------------------------------------------------------------ + # Outbound: requests this side sends and correlates against responses. + # ------------------------------------------------------------------ + + def allocate_id(self) -> int: + """Mint the next dispatcher-owned request id (monotonic, starts at 1).""" + self._next_id += 1 + return self._next_id + + def _reserve_id(self, supplied: RequestId | None) -> tuple[RequestId, RequestId]: + """Pick the wire id and the pending-table key for one outbound request. + + A caller-supplied id is used verbatim on the wire and coerced for the + key; a collision with an in-flight key raises `ValueError`. Otherwise + a fresh id is minted past any key a supplied id occupies: the collision + error is reserved for the caller who actually chose the id. + """ + if supplied is not None: + key = coerce_request_id(supplied) + if key in self.pending: + raise ValueError(f"request id {supplied!r} is already in flight") + return supplied, key + request_id = self.allocate_id() + while request_id in self.pending: + request_id = self.allocate_id() + return request_id, request_id + + async def call( + self, + method: str, + params: Mapping[str, Any] | None, + opts: CallOptions | None, + *, + write_request: Callable[[JSONRPCRequest], Awaitable[None]], + send_cancel: Callable[[RequestId, str], Awaitable[None]], + cancel_on_abandon: bool, + ) -> dict[str, Any]: + """Send one outbound request and await its correlated response. + + `write_request` puts the built `JSONRPCRequest` on the wire (raising + `anyio.BrokenResourceError` / `ClosedResourceError` means the channel + is gone). `send_cancel(request_id, reason)` emits the courtesy + `notifications/cancelled` on the abandon paths when + `cancel_on_abandon` is set; it must swallow its own write failures. + + Raises: + MCPError: Peer error response; `REQUEST_TIMEOUT` if + `opts["timeout"]` elapsed; `CONNECTION_CLOSED` if closed or + the write channel was torn down. + ValueError: `opts["request_id"]` collides with an in-flight id. + """ + # Post-close sends get the same CONNECTION_CLOSED contract as in-flight waiters. + if self._closed: + raise MCPError(code=CONNECTION_CLOSED, message="Connection closed") + opts = opts or {} + request_id, pending_key = self._reserve_id(opts.get("request_id")) + out_params = dict(params) if params is not None else {} + out_meta = dict(out_params.get("_meta") or {}) + on_progress = opts.get("on_progress") + if on_progress is not None: + # The request id doubles as the progress token, so `pending[token]` finds `on_progress` directly. + out_meta["progressToken"] = request_id + out_params["_meta"] = out_meta + + # buffer=1: a close signal can arrive before the waiter parks in receive(); + # a WouldBlock later just means the waiter already has its one outcome. + send, receive = anyio.create_memory_object_stream[Outcome](1) + pending = Pending(send=send, receive=receive, on_progress=on_progress) + self.pending[pending_key] = pending + + # Spec MUST: only previously-issued requests may be cancelled. A write + # interrupted by cancellation may still have delivered (a memory-stream + # send can hand its item to the receiver and still raise), so a started + # write counts as issued: the peer ignores a cancel for an id it never + # saw, while skipping it would leak a delivered request's handler. + request_write_started = False + timeout_armed = False + + target = out_params.get("name") + span_name = f"MCP send {method}{f' {target}' if isinstance(target, str) else ''}" + # TODO(maxisbey): move the otel span + inject into an outbound + # middleware once that seam exists; the correlator should not own otel. + try: + with otel_span( + span_name, + kind=SpanKind.CLIENT, + attributes={"mcp.method.name": method, "jsonrpc.request.id": str(request_id)}, + ): + # SEP-414: inject W3C trace context; `_meta` stays on the wire even with a no-op tracer. + inject_trace_context(out_meta) + msg = JSONRPCRequest(jsonrpc="2.0", id=request_id, method=method, params=out_params) + # Surface a pre-existing cancellation while the request provably + # never started; past this point a cancelled write counts as issued. + await anyio.lowlevel.checkpoint_if_cancelled() + request_write_started = True + try: + await write_request(msg) + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + # Channel tore down before its owner noticed EOF; surface the documented contract. + raise MCPError(code=CONNECTION_CLOSED, message="Connection closed") from None + with anyio.fail_after(opts.get("timeout")): + timeout_armed = True + outcome = await receive.receive() + except TimeoutError: + if not timeout_armed: + # `fail_after` arms only after the write, so this TimeoutError is the + # channel's own bounded send() failing - a transport error, not + # `opts["timeout"]` elapsing. Propagate it raw (v1 kept the write + # outside the timeout-catching try and did the same). + raise + # Courtesy cancel (spec-recommended) so the peer stops work; + # unshielded so an outer caller cancellation can still interrupt the write. + if cancel_on_abandon: + await final_write( + partial(send_cancel, request_id, f"timed out after {opts.get('timeout')}s"), + shield=False, + timeout=_ABANDON_WRITE_TIMEOUT, + describe=f"courtesy cancel for timed-out request {request_id!r}", + ) + raise MCPError(code=REQUEST_TIMEOUT, message=f"Request {method!r} timed out") from None + except anyio.get_cancelled_exc_class(): + # Caller cancelled: bare awaits re-raise here, so the shielded helper + # lets the courtesy cancel go out before we propagate. + if cancel_on_abandon and request_write_started: + await final_write( + partial(send_cancel, request_id, "caller cancelled"), + shield=True, + timeout=_ABANDON_WRITE_TIMEOUT, + describe=f"courtesy cancel for caller-cancelled request {request_id!r}", + ) + raise + finally: + # Remove the waiter on every path so a late response is dropped, not leaked. + self.pending.pop(pending_key, None) + send.close() + receive.close() + + if isinstance(outcome, ErrorData): + raise MCPError(code=outcome.code, message=outcome.message, data=outcome.data) + return outcome + + def resolve(self, request_id: RequestId | None, outcome: Outcome) -> None: + """Deliver `outcome` to the waiter for `request_id`; unknown/late ids are dropped.""" + pending = self.pending.get(coerce_request_id(request_id)) if request_id is not None else None + if pending is None: + logger.debug("dropping response for unknown/late request id %r", request_id) + return + try: + pending.send.send_nowait(outcome) + except (anyio.WouldBlock, anyio.BrokenResourceError, anyio.ClosedResourceError): + logger.debug("waiter for request id %r already gone", request_id) + + def progress_callback( + self, params: Mapping[str, Any] | None + ) -> tuple[ProgressFnT, float, float | None, str | None] | None: + """Match a `notifications/progress` body against the pending outbound requests. + + Returns the shielded callback and its coerced arguments when the + token names one of our requests that asked for progress, else `None`. + `bool` is rejected everywhere it would alias an int/float. + """ + match params: + case {"progressToken": str() | int() as token, "progress": int() | float() as progress} if ( + not isinstance(token, bool) + and not isinstance(progress, bool) + and (pending := self.pending.get(coerce_request_id(token))) is not None + and pending.on_progress is not None + ): + total = params.get("total") + message = params.get("message") + return ( + _shielded_progress(pending.on_progress), + float(progress), + float(total) if isinstance(total, int | float) else None, + message if isinstance(message, str) else None, + ) + case _: + return None + + def fan_out_closed(self) -> None: + """Wake every pending outbound waiter with `CONNECTION_CLOSED`. + + Synchronous: callers may be inside a cancelled scope. Idempotent. + """ + closed = ErrorData(code=CONNECTION_CLOSED, message="Connection closed") + for pending in self.pending.values(): + try: + pending.send.send_nowait(closed) + except (anyio.WouldBlock, anyio.BrokenResourceError, anyio.ClosedResourceError): + pass + self.pending.clear() + + # ------------------------------------------------------------------ + # Inbound: requests the peer sends and this side handles. + # ------------------------------------------------------------------ + + def enter_inbound(self, request_id: RequestId, scope: anyio.CancelScope, dctx: DctxT) -> None: + """Register an inbound request before its handler runs, so peer cancels can find it. + + Duplicate ids blind-overwrite (v1/TS parity); the identity guard in + `serve_inbound` keeps a superseded entry from evicting its successor. + """ + # TODO(maxisbey): duplicate ids blind-overwrite (v1/TS parity); revisit + # rejecting with INVALID_REQUEST. Key coerced so a stringified + # `notifications/cancelled` id still correlates. + self.in_flight[coerce_request_id(request_id)] = InFlight(scope=scope, dctx=dctx) + + def peer_cancel(self, request_id: RequestId | None, *, interrupt: bool) -> bool: + """Apply a peer's `notifications/cancelled` for `request_id`. + + Sets the handler's `cancel_requested` event and, when `interrupt`, + cancels its scope. Returns whether a matching in-flight request was found. + """ + if request_id is None: + return False + in_flight = self.in_flight.get(coerce_request_id(request_id)) + if in_flight is None: + return False + in_flight.dctx.cancel_requested.set() + if interrupt: + in_flight.scope.cancel() + return True + + def cancel_all_inbound(self) -> None: + """Cancel every in-flight handler's scope (shutdown/termination).""" + for entry in list(self.in_flight.values()): + entry.scope.cancel() + + async def serve_inbound( + self, + request_id: RequestId, + dctx: DctxT, + scope: anyio.CancelScope, + run: Callable[[], Awaitable[dict[str, Any]]], + *, + write_result: Callable[[dict[str, Any]], Awaitable[None]], + write_error: Callable[[ErrorData], Awaitable[None]], + settle_unanswered: Callable[[], Awaitable[None]] | None = None, + raise_handler_exceptions: bool = False, + ) -> None: + """Run one registered inbound request and write at most one terminal outcome. + + The single exception-to-wire boundary for inbound requests. `run` is + the handler invocation; `write_result`/`write_error` put the terminal + response on whatever channel the transport uses (they must not raise + for a torn-down channel). A request the peer cancelled is never + answered (spec: MUST NOT send further messages for it): its result or + error is dropped and it settles through `settle_unanswered` - the + transport's hook for a wire that must still end the request another + way. The caller registers `(request_id, scope, dctx)` via + `enter_inbound` *before* scheduling this, so a peer cancel that races + the handler's start still lands. + """ + answer_write_started = False + handler_failure: BaseException | None = None # re-raised once the request settles + try: + with scope: + try: + result = await run() + finally: + # Close the back-channel and drop from `in_flight`; no checkpoint + # since handler return, so a peer cancel can't interleave. + # Identity guard: don't evict a duplicate id's newer entry. + dctx.close() + key = coerce_request_id(request_id) + if (entry := self.in_flight.get(key)) is not None and entry.dctx is dctx: + del self.in_flight[key] + if not dctx.cancel_requested.is_set(): + # A write interrupted by cancellation may still have delivered + # (a memory-stream send can hand its item to the receiver and + # still raise), so a started answer write counts as sent below: + # peers drop late responses, while a second answer for one id + # would break JSON-RPC. + answer_write_started = True + await write_result(result) + except anyio.get_cancelled_exc_class(): + # Shutdown: answer the request so the peer isn't left waiting - unless + # an answer write already started (it may have reached the channel; + # prefer possibly-zero answers over possibly-two), or the peer already + # cancelled it and stopped waiting. The shielded helper is needed + # because bare awaits re-raise here. + if not answer_write_started and not dctx.cancel_requested.is_set(): + await final_write( + partial(write_error, ErrorData(code=CONNECTION_CLOSED, message="Connection closed")), + shield=True, + timeout=_SHUTDOWN_WRITE_TIMEOUT, + describe=f"shutdown error response for request {request_id!r}", + ) + raise + except Exception as e: + error = handler_exception_to_error_data(e) + if error is None: + logger.exception("handler for request %r raised", request_id) + # TODO(L58): code=0 pins existing-server compat; JSON-RPC says + # INTERNAL_ERROR. Revisit per the suite's divergence entry. + error = ErrorData(code=0, message=str(e)) + if raise_handler_exceptions: + handler_failure = e + # A cancel silences only the wire; the failure stays as visible as before. + if not dctx.cancel_requested.is_set(): + answer_write_started = True + await write_error(error) + # The one place a cancelled request settles: the handler is done (any + # mode) with nothing written. A peer-interrupt cancel is absorbed at + # scope __exit__ and lands here too. + if not answer_write_started and settle_unanswered is not None: + try: + await settle_unanswered() + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + logger.debug("on_request_unanswered dropped: connection closing") + except Exception: + logger.exception("on_request_unanswered hook raised") + if handler_failure is not None: + raise handler_failure + # No `in_flight` pop here: the inner finally covers every path, and a late pop could evict a reused id. diff --git a/src/mcp/shared/jsonrpc_dispatcher.py b/src/mcp/shared/jsonrpc_dispatcher.py index 0d9467ffeb..4ff7afb260 100644 --- a/src/mcp/shared/jsonrpc_dispatcher.py +++ b/src/mcp/shared/jsonrpc_dispatcher.py @@ -1,8 +1,10 @@ """JSON-RPC `Dispatcher` over the `SessionMessage` stream contract all transports speak. -Owns request-id correlation, the receive loop, per-request task isolation, -cancellation/progress wiring, and the single exception-to-wire boundary; -methods and params are otherwise opaque strings and dicts. +Owns the receive loop and per-request task isolation over a duplex stream +pair; request-id correlation, cancellation/progress wiring, and the single +exception-to-wire boundary live in the shared `RequestCorrelator` so the +streamable-HTTP transport (which has no stream pair) applies the same +semantics. Methods and params are otherwise opaque strings and dicts. """ from __future__ import annotations @@ -16,13 +18,8 @@ import anyio import anyio.abc -import anyio.lowlevel -from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from mcp_types import ( - CONNECTION_CLOSED, INTERNAL_ERROR, - INVALID_PARAMS, - REQUEST_TIMEOUT, ErrorData, JSONRPCError, JSONRPCMessage, @@ -32,12 +29,16 @@ ProgressToken, RequestId, ) -from opentelemetry.trace import SpanKind -from pydantic import ValidationError from typing_extensions import TypeVar from mcp.shared._compat import resync_tracer -from mcp.shared._otel import inject_trace_context, otel_span +from mcp.shared._correlation import ( + InFlight, + Outcome, + Pending, + RequestCorrelator, + handler_exception_to_error_data, +) from mcp.shared._stream_protocols import ReadStream, WriteStream from mcp.shared.dispatcher import ( CallOptions, @@ -46,12 +47,10 @@ OnNotify, OnNotifyIntercept, OnRequest, - ProgressFnT, as_request_id, - coerce_request_id, run_notify_intercept, ) -from mcp.shared.exceptions import MCPError, NoBackChannelError +from mcp.shared.exceptions import NoBackChannelError from mcp.shared.message import ( ClientMessageMetadata, MessageMetadata, @@ -69,13 +68,6 @@ logger = logging.getLogger(__name__) -_ABANDON_WRITE_TIMEOUT: float = 5 -"""Bound for courtesy-cancel writes on the abandon paths; the caller-cancel -arm shields its write, so a wedged transport would otherwise hang it uncancellably.""" - -_SHUTDOWN_WRITE_TIMEOUT: float = 1 -"""Tighter bound for the shutdown-arm error write so a wedged transport can't hold session close.""" - TransportT = TypeVar("TransportT", bound=TransportContext, default=TransportContext) PeerCancelMode = Literal["interrupt", "signal"] @@ -84,22 +76,8 @@ handler run to completion. Either way the cancelled request is never answered - the handler's eventual result or error is dropped, not written.""" - -def handler_exception_to_error_data(exc: BaseException) -> ErrorData | None: - """Map a handler-raised exception to its wire `ErrorData`. - - The two rungs every dispatcher shares: an `MCPError` carries its own - `ErrorData`; a pydantic `ValidationError` is the spec's INVALID_PARAMS - with empty ``data`` (no pydantic text on the wire). Returns ``None`` for - any other exception so each caller applies its own catch-all - - `JSONRPCDispatcher` currently pins ``code=0`` for v1 compat, - the modern HTTP entry uses `INTERNAL_ERROR`. - """ - if isinstance(exc, MCPError): - return exc.error - if isinstance(exc, ValidationError): - return ErrorData(code=INVALID_PARAMS, message="Invalid request parameters", data="") - return None +_Pending = Pending +"""Outbound-waiter record; owned by `RequestCorrelator` (aliased here for white-box tests).""" def progress_token_from_params(params: Mapping[str, Any] | None) -> ProgressToken | None: @@ -116,23 +94,6 @@ def cancelled_request_id_from_params(params: Mapping[str, Any] | None) -> Reques return as_request_id((params or {}).get("requestId")) -@dataclass(slots=True) -class _Pending: - """An outbound request awaiting its response.""" - - send: MemoryObjectSendStream[dict[str, Any] | ErrorData] - receive: MemoryObjectReceiveStream[dict[str, Any] | ErrorData] - on_progress: ProgressFnT | None = None - - -@dataclass(slots=True) -class _InFlight(Generic[TransportT]): - """An inbound request currently being handled.""" - - scope: anyio.CancelScope - dctx: _JSONRPCDispatchContext[TransportT] - - @dataclass class _JSONRPCDispatchContext(Generic[TransportT]): """Concrete `DispatchContext` produced for each inbound JSON-RPC message.""" @@ -188,20 +149,8 @@ def _default_transport_builder(_meta: MessageMetadata) -> TransportContext: return TransportContext(kind="jsonrpc", can_send_request=True) -def _shielded_progress(fn: ProgressFnT) -> ProgressFnT: - """Wrap a user progress callback so an exception can't cancel the dispatcher's task group.""" - - async def _wrapped(progress: float, total: float | None, message: str | None) -> None: - try: - await fn(progress, total, message) - except Exception: - logger.exception("progress callback raised") - - return _wrapped - - def _contained_notify(fn: OnNotify) -> OnNotify: - """Wrap a notification handler so it can't crash the dispatcher (same boundary as `_shielded_progress`).""" + """Wrap a notification handler so it can't crash the dispatcher's task group.""" async def _wrapped(dctx: DispatchContext[TransportContext], method: str, params: Mapping[str, Any] | None) -> None: try: @@ -295,13 +244,14 @@ def __init__( bind it after the dispatcher is built (e.g. ``ClientSession`` routing into ``message_handler``); only consulted inside ``run()`` so pre-enter assignment is safe.""" - self._next_id = 0 - self._pending: dict[RequestId, _Pending] = {} - self._in_flight: dict[RequestId, _InFlight[TransportT]] = {} + # The correlation kernel owns the pending/in-flight tables; the + # aliases keep the historical private names white-box tests read. + self._corr: RequestCorrelator[_JSONRPCDispatchContext[TransportT]] = RequestCorrelator() + self._pending: dict[RequestId, Pending] = self._corr.pending + self._in_flight: dict[RequestId, InFlight[_JSONRPCDispatchContext[TransportT]]] = self._corr.in_flight self._on_notify_intercept: OnNotifyIntercept | None = None self._tg: anyio.abc.TaskGroup | None = None self._running = False - self._closed = False async def send_raw_request( self, @@ -322,118 +272,19 @@ async def send_raw_request( transport closed or the dispatcher shut down. RuntimeError: Called before `run()`. """ - # Post-close sends get the same CONNECTION_CLOSED contract as in-flight waiters. - if self._closed: - raise MCPError(code=CONNECTION_CLOSED, message="Connection closed") - if not self._running: + # Post-close sends get the same CONNECTION_CLOSED contract as in-flight + # waiters (raised by the correlator); only a never-run dispatcher is a usage error. + if not self._running and not self._corr.closed: raise RuntimeError("JSONRPCDispatcher.send_raw_request called before run()") - opts = opts or {} - supplied_id = opts.get("request_id") - if supplied_id is not None: - request_id: RequestId = supplied_id - # The pending key gets the same coercion `_resolve_pending` applies - # to inbound response ids, so a supplied "7" still correlates - # whether the peer echoes "7" or 7. The wire id stays verbatim. - pending_key = coerce_request_id(request_id) - if pending_key in self._pending: - raise ValueError(f"request id {request_id!r} is already in flight") - else: - # Mint past any key a supplied id occupies: the collision error is - # reserved for the caller who actually chose the id. - request_id = self._allocate_id() - while request_id in self._pending: - request_id = self._allocate_id() - pending_key = request_id - out_params = dict(params) if params is not None else {} - out_meta = dict(out_params.get("_meta") or {}) - on_progress = opts.get("on_progress") - if on_progress is not None: - # The request id doubles as the progress token, so `_pending[token]` finds `on_progress` directly. - out_meta["progressToken"] = request_id - out_params["_meta"] = out_meta - - # buffer=1: a close signal can arrive before the waiter parks in receive(); - # a WouldBlock later just means the waiter already has its one outcome. - send, receive = anyio.create_memory_object_stream[dict[str, Any] | ErrorData](1) - pending = _Pending(send=send, receive=receive, on_progress=on_progress) - self._pending[pending_key] = pending - plan = _plan_outbound(_related_request_id, opts) - # Spec MUST: only previously-issued requests may be cancelled. A write - # interrupted by cancellation may still have delivered (a memory-stream - # send can hand its item to the receiver and still raise), so a started - # write counts as issued: the peer ignores a cancel for an id it never - # saw, while skipping it would leak a delivered request's handler. - request_write_started = False - timeout_armed = False - - target = out_params.get("name") - span_name = f"MCP send {method}{f' {target}' if isinstance(target, str) else ''}" - # TODO(maxisbey): move the otel span + inject into an outbound - # middleware once that seam exists; the dispatcher should not own otel. - try: - with otel_span( - span_name, - kind=SpanKind.CLIENT, - attributes={"mcp.method.name": method, "jsonrpc.request.id": str(request_id)}, - ): - # SEP-414: inject W3C trace context; `_meta` stays on the wire even with a no-op tracer. - inject_trace_context(out_meta) - msg = JSONRPCRequest(jsonrpc="2.0", id=request_id, method=method, params=out_params) - # Surface a pre-existing cancellation while the request provably - # never started; past this point a cancelled write counts as issued. - await anyio.lowlevel.checkpoint_if_cancelled() - request_write_started = True - try: - await self._write(msg, plan.metadata) - except (anyio.BrokenResourceError, anyio.ClosedResourceError): - # Transport tore down before run() noticed EOF; surface the documented contract. - raise MCPError(code=CONNECTION_CLOSED, message="Connection closed") from None - with anyio.fail_after(opts.get("timeout")): - timeout_armed = True - outcome = await receive.receive() - except TimeoutError: - if not timeout_armed: - # `fail_after` arms only after the write, so this TimeoutError is the - # transport's own bounded send() failing - a transport error, not - # `opts["timeout"]` elapsing. Propagate it raw (v1 kept the write - # outside the timeout-catching try and did the same). - raise - # Courtesy cancel (spec-recommended, new vs v1) so the peer stops work; - # unshielded so an outer caller cancellation can still interrupt the write. - if plan.cancel_on_abandon: - await self._final_write( - partial( - self._cancel_outbound, - request_id, - f"timed out after {opts.get('timeout')}s", - _related_request_id, - ), - shield=False, - timeout=_ABANDON_WRITE_TIMEOUT, - describe=f"courtesy cancel for timed-out request {request_id!r}", - ) - raise MCPError(code=REQUEST_TIMEOUT, message=f"Request {method!r} timed out") from None - except anyio.get_cancelled_exc_class(): - # Caller cancelled: bare awaits re-raise here, so the shielded helper - # lets the courtesy cancel go out before we propagate. - if plan.cancel_on_abandon and request_write_started: - await self._final_write( - partial(self._cancel_outbound, request_id, "caller cancelled", _related_request_id), - shield=True, - timeout=_ABANDON_WRITE_TIMEOUT, - describe=f"courtesy cancel for caller-cancelled request {request_id!r}", - ) - raise - finally: - # Remove the waiter on every path so a late response is dropped, not leaked. - self._pending.pop(pending_key, None) - send.close() - receive.close() - - if isinstance(outcome, ErrorData): - raise MCPError(code=outcome.code, message=outcome.message, data=outcome.data) - return outcome + return await self._corr.call( + method, + params, + opts, + write_request=partial(self._write, metadata=plan.metadata), + send_cancel=partial(self._cancel_outbound, related_request_id=_related_request_id), + cancel_on_abandon=plan.cancel_on_abandon, + ) async def notify( self, @@ -449,7 +300,7 @@ async def notify( torn-down transport drops the notification with a debug log instead of raising (same policy as the response writes and `ctx.notify`). """ - if self._closed: + if self._corr.closed: logger.debug("dropped %s: dispatcher closed", method) return # Leave `params` unset when None: with `exclude_unset=True` an explicit @@ -500,18 +351,16 @@ async def run( logger.debug("read stream closed by transport; treating as EOF") # EOF: wake blocked `send_raw_request` waiters with CONNECTION_CLOSED. self._running = False - self._closed = True - self._fan_out_closed() + self._corr.close() finally: # Cancel in-flight handlers; otherwise the task-group join # waits on handlers whose callers are already gone. tg.cancel_scope.cancel() finally: - # Covers cancel/crash paths that skip the inline fan-out; idempotent. + # Covers cancel/crash paths that skip the inline close; idempotent. self._running = False - self._closed = True self._tg = None - self._fan_out_closed() + self._corr.close() await resync_tracer() async def _dispatch( @@ -576,10 +425,7 @@ async def _dispatch_request( _progress_token=progress_token, ) scope = anyio.CancelScope() - # TODO(maxisbey): duplicate ids blind-overwrite (v1/TS parity); revisit - # rejecting with INVALID_REQUEST. Key coerced so a stringified - # `notifications/cancelled` id still correlates. - self._in_flight[coerce_request_id(req.id)] = _InFlight(scope=scope, dctx=dctx) + self._corr.enter_inbound(req.id, scope, dctx) if req.method in self._inline_methods: # Spawn so `sender_ctx` applies, but park the read loop until the # handler returns - that's the inline ordering guarantee. @@ -606,36 +452,21 @@ def _dispatch_notification( """Route one inbound notification. `notifications/cancelled` and `notifications/progress` are intercepted - here (they correlate against the `_in_flight`/`_pending` tables this - layer owns) and still teed to `on_notify` afterwards. The caller's + here (they correlate against the correlator's in-flight/pending + tables) and still teed to `on_notify` afterwards. The caller's `on_notify_intercept` then runs in receive order; only unconsumed notifications reach the spawned `on_notify`. """ if msg.method == "notifications/cancelled": - rid = cancelled_request_id_from_params(msg.params) - if rid is not None and (in_flight := self._in_flight.get(coerce_request_id(rid))) is not None: - in_flight.dctx.cancel_requested.set() - if self._peer_cancel_mode == "interrupt": - in_flight.scope.cancel() + self._corr.peer_cancel( + cancelled_request_id_from_params(msg.params), + interrupt=self._peer_cancel_mode == "interrupt", + ) elif msg.method == "notifications/progress": - match msg.params: - case {"progressToken": str() | int() as token, "progress": int() | float() as progress} if ( - not isinstance(token, bool) - and not isinstance(progress, bool) - and (pending := self._pending.get(coerce_request_id(token))) is not None - and pending.on_progress is not None - ): - total = msg.params.get("total") - message = msg.params.get("message") - self._spawn( - _shielded_progress(pending.on_progress), - float(progress), - float(total) if isinstance(total, int | float) else None, - message if isinstance(message, str) else None, - sender_ctx=sender_ctx, - ) - case _: - pass + delivery = self._corr.progress_callback(msg.params) + if delivery is not None: + fn, progress, total, message = delivery + self._spawn(fn, progress, total, message, sender_ctx=sender_ctx) if run_notify_intercept(self._on_notify_intercept, msg.method, msg.params): return try: @@ -649,15 +480,8 @@ def _dispatch_notification( ) self._spawn(_contained_notify(on_notify), dctx, msg.method, msg.params, sender_ctx=sender_ctx) - def _resolve_pending(self, request_id: RequestId | None, outcome: dict[str, Any] | ErrorData) -> None: - pending = self._pending.get(coerce_request_id(request_id)) if request_id is not None else None - if pending is None: - logger.debug("dropping response for unknown/late request id %r", request_id) - return - try: - pending.send.send_nowait(outcome) - except (anyio.WouldBlock, anyio.BrokenResourceError, anyio.ClosedResourceError): - logger.debug("waiter for request id %r already gone", request_id) + def _resolve_pending(self, request_id: RequestId | None, outcome: Outcome) -> None: + self._corr.resolve(request_id, outcome) def _spawn( self, @@ -677,17 +501,8 @@ def _spawn( self._tg.start_soon(fn, *args) def _fan_out_closed(self) -> None: - """Wake every pending `send_raw_request` waiter with `CONNECTION_CLOSED`. - - Synchronous: callers may be inside a cancelled scope. Idempotent. - """ - closed = ErrorData(code=CONNECTION_CLOSED, message="Connection closed") - for pending in self._pending.values(): - try: - pending.send.send_nowait(closed) - except (anyio.WouldBlock, anyio.BrokenResourceError, anyio.ClosedResourceError): - pass - self._pending.clear() + """Wake every pending `send_raw_request` waiter with `CONNECTION_CLOSED`. Idempotent.""" + self._corr.fan_out_closed() async def _handle_request( self, @@ -698,72 +513,22 @@ async def _handle_request( ) -> None: """Run `on_request` for one inbound request and write its response. - The single exception-to-wire boundary: handler exceptions become - `JSONRPCError` here. A request the peer cancelled is never answered - (spec: MUST NOT send further messages for it) - it settles unanswered - instead, and `_settle_unanswered` tells the transport. + The exception-to-wire policy lives in `RequestCorrelator.serve_inbound`; + this only binds the wire writes for a stream-pair transport. A request + the peer cancelled is never answered (spec: MUST NOT send further + messages for it) - it settles unanswered instead, and `_settle_unanswered` + tells the transport. """ - answer_write_started = False - handler_failure: BaseException | None = None # re-raised once the request settles - try: - with scope: - try: - result = await on_request(dctx, req.method, req.params) - finally: - # Close the back-channel and drop from `_in_flight`; no checkpoint - # since handler return, so a peer cancel can't interleave. - # Identity guard: don't evict a duplicate id's newer entry. - dctx.close() - key = coerce_request_id(req.id) - if (entry := self._in_flight.get(key)) is not None and entry.dctx is dctx: - del self._in_flight[key] - if not dctx.cancel_requested.is_set(): - # A write interrupted by cancellation may still have delivered - # (a memory-stream send can hand its item to the receiver and - # still raise), so a started answer write counts as sent below: - # peers drop late responses, while a second answer for one id - # would break JSON-RPC. - answer_write_started = True - await self._write_result(req.id, result) - except anyio.get_cancelled_exc_class(): - # Shutdown: answer the request so the peer isn't left waiting - unless - # an answer write already started (it may have reached the transport; - # prefer possibly-zero answers over possibly-two), or the peer already - # cancelled it and stopped waiting. The shielded helper is needed - # because bare awaits re-raise here. - if not answer_write_started and not dctx.cancel_requested.is_set(): - await self._final_write( - partial(self._write_error, req.id, ErrorData(code=CONNECTION_CLOSED, message="Connection closed")), - shield=True, - timeout=_SHUTDOWN_WRITE_TIMEOUT, - describe=f"shutdown error response for request {req.id!r}", - ) - raise - except Exception as e: - error = handler_exception_to_error_data(e) - if error is None: - logger.exception("handler for %r raised", req.method) - # TODO(L58): code=0 pins existing-server compat; JSON-RPC says - # INTERNAL_ERROR. Revisit per the suite's divergence entry. - error = ErrorData(code=0, message=str(e)) - if self._raise_handler_exceptions: - handler_failure = e - # A cancel silences only the wire; the failure stays as visible as before. - if not dctx.cancel_requested.is_set(): - answer_write_started = True - await self._write_error(req.id, error) - # The one place a cancelled request settles: the handler is done (any - # mode) with nothing written. A peer-interrupt cancel is absorbed at - # scope __exit__ and lands here too. - if not answer_write_started: - await self._settle_unanswered(dctx) - if handler_failure is not None: - raise handler_failure - # No `_in_flight` pop here: the inner finally covers every path, and a late pop could evict a reused id. - - def _allocate_id(self) -> int: - self._next_id += 1 - return self._next_id + await self._corr.serve_inbound( + req.id, + dctx, + scope, + partial(on_request, dctx, req.method, req.params), + write_result=partial(self._write_result, req.id), + write_error=partial(self._write_error, req.id), + settle_unanswered=partial(self._settle_unanswered, dctx), + raise_handler_exceptions=self._raise_handler_exceptions, + ) async def _write(self, message: JSONRPCMessage, metadata: MessageMetadata = None) -> None: await self._write_stream.send(SessionMessage(message=message, metadata=metadata)) @@ -784,37 +549,13 @@ async def _settle_unanswered(self, dctx: _JSONRPCDispatchContext[TransportT]) -> """Run the transport's `on_request_unanswered` hook: this request settled with no response. The dispatcher writes nothing for it; a transport whose wire must still - end the request (2025-era streamable HTTP) does so from this hook. A - raising hook is contained here, like the other callback boundaries. + end the request (2025-era streamable HTTP) does so from this hook. + `RequestCorrelator.serve_inbound` invokes and contains it. """ metadata = dctx.message_metadata if not isinstance(metadata, ServerMessageMetadata) or metadata.on_request_unanswered is None: return - try: - await metadata.on_request_unanswered() - except (anyio.BrokenResourceError, anyio.ClosedResourceError): - logger.debug("on_request_unanswered dropped: connection closing") - except Exception: - logger.exception("on_request_unanswered hook raised") - - async def _final_write( - self, - write: Callable[[], Awaitable[None]], - *, - shield: bool, - timeout: float, - describe: str, - ) -> None: - """Attempt one last write under the shared abandon/teardown policy. - - `shield=True` is for arms already inside a cancelled scope (a bare - `await` would re-raise); the bound keeps a wedged transport write - from becoming an uncancellable hang. - """ - with anyio.move_on_after(timeout, shield=shield) as scope: - await write() - if scope.cancelled_caught: - logger.warning("%s gave up: transport write blocked", describe) + await metadata.on_request_unanswered() async def _cancel_outbound(self, request_id: RequestId, reason: str, related_request_id: RequestId | None) -> None: # Thread `related_request_id` so streamable HTTP routes the cancel onto diff --git a/tests/interaction/README.md b/tests/interaction/README.md index 3060a240c1..17a4a7c332 100644 --- a/tests/interaction/README.md +++ b/tests/interaction/README.md @@ -279,9 +279,10 @@ this hits any test that must run statements after a `ClientSession`/`streamable_ but still inside an outer `async with`, and no restructure can avoid it. A handful of `# pragma: lax no cover` markers in `src/` cover teardown exception handlers whose -execution is timing-dependent under the in-process HTTP bridge — the POST-stream and -stateless-session `except Exception` handlers in `server/streamable_http*.py` and the -`_terminated` check in `message_router`. `strict-no-cover` does not check `lax` lines; do not +execution is timing-dependent under the in-process HTTP bridge — the `except Exception` arms +around the SSE-response runner (`_run_sse_response`) and the replay entry path in +`server/streamable_http.py`. +`strict-no-cover` does not check `lax` lines; do not promote them to strict `no cover` without first making the teardown ordering deterministic. The suite also relies on a one-line `src/mcp/server/sse.py` fix (`sse_stream_reader.aclose()`) that closes a stream the SSE leg would otherwise leak. diff --git a/tests/server/test_runner.py b/tests/server/test_runner.py index eb212dafdb..7ef3f404a3 100644 --- a/tests/server/test_runner.py +++ b/tests/server/test_runner.py @@ -53,6 +53,7 @@ ) import mcp.server.runner +from mcp.client.session import ClientSession from mcp.server.caching import CacheHint from mcp.server.connection import Connection, NotifyOnlyOutbound from mcp.server.context import ServerRequestContext @@ -67,6 +68,7 @@ aclose_shielded, serve_connection, serve_dual_era_loop, + serve_loop, serve_one, ) from mcp.server.session import ServerSession @@ -75,6 +77,7 @@ from mcp.shared.dispatcher import CallOptions from mcp.shared.exceptions import MCPError, NoBackChannelError from mcp.shared.jsonrpc_dispatcher import JSONRPCDispatcher +from mcp.shared.memory import create_client_server_memory_streams from mcp.shared.message import MessageMetadata, SessionMessage from mcp.shared.peer import dump_params from mcp.shared.transport_context import TransportContext @@ -2056,3 +2059,22 @@ async def test_dual_era_client_propagates_body_exception_unwrapped(server: SrvT) with pytest.raises(RuntimeError, match="boom"): async with dual_era_client(server): raise RuntimeError("boom") + + +@pytest.mark.anyio +async def test_serve_loop_serves_a_handshake_connection_over_a_stream_pair(server: SrvT) -> None: + """`serve_loop`, the loop-mode driver for transports that own their own lifespan, round-trips + a handshake and a request over a duplex stream pair and returns when the channel closes.""" + async with create_client_server_memory_streams() as ((client_read, client_write), (server_read, server_write)): + with anyio.fail_after(5): + async with anyio.create_task_group() as tg: # pragma: no branch + tg.start_soon( + partial(serve_loop, server, server_read, server_write, lifespan_state={}, session_id="loop-1") + ) + async with ClientSession(client_read, client_write) as session: + initialized = await session.initialize() + tools = await session.list_tools() + assert initialized.server_info.name == "test-server" + assert [tool.name for tool in tools.tools] == ["t"] + # Closing the client's write side EOFs the loop and lets it return. + await client_write.aclose() diff --git a/tests/server/test_streamable_http_manager.py b/tests/server/test_streamable_http_manager.py index 70440d9d03..17d216c887 100644 --- a/tests/server/test_streamable_http_manager.py +++ b/tests/server/test_streamable_http_manager.py @@ -9,6 +9,7 @@ import anyio import httpx2 import pytest +from anyio.abc import TaskStatus from mcp_types import INVALID_REQUEST, ListToolsResult, PaginatedRequestParams from starlette.types import Message, Receive, Scope, Send @@ -252,10 +253,14 @@ async def running_manager(): async def test_stateful_session_cleanup_on_graceful_exit(running_manager: tuple[StreamableHTTPSessionManager, Server]): manager, _app = running_manager - # The manager's `run_server` task drives `serve_loop` directly (the manager - # owns lifespan); patch that seam so the loop returns immediately and we - # can observe the cleanup that follows. - mock_serve = AsyncMock(return_value=None) + # The manager's `run_server` task drives the transport's session task + # (`run()`); patch that seam so it returns immediately and we can observe + # the cleanup that follows. + run_calls: list[None] = [] + + async def mock_run(self: StreamableHTTPServerTransport, *, task_status: TaskStatus[None]) -> None: + run_calls.append(None) + task_status.started() sent_messages: list[Message] = [] @@ -273,7 +278,7 @@ async def mock_receive(): return {"type": "http.request", "body": b"", "more_body": False} # Trigger session creation - with patch("mcp.server.streamable_http_manager.serve_loop", mock_serve): + with patch.object(StreamableHTTPServerTransport, "run", mock_run): await manager.handle_request(scope, mock_receive, mock_send) # Extract session ID from response headers @@ -289,9 +294,9 @@ async def mock_receive(): assert session_id is not None, "Session ID not found in response headers" - mock_serve.assert_called_once() + assert len(run_calls) == 1 - # At this point, mock_serve has completed, and the finally block in + # At this point, mock_run has completed, and the finally block in # StreamableHTTPSessionManager's run_server should have executed. # To ensure the task spawned by handle_request finishes and cleanup occurs: @@ -308,7 +313,12 @@ async def mock_receive(): async def test_stateful_session_cleanup_on_exception(running_manager: tuple[StreamableHTTPSessionManager, Server]): manager, _app = running_manager - mock_serve = AsyncMock(side_effect=TestException("Simulated crash")) + run_calls: list[None] = [] + + async def mock_run(self: StreamableHTTPServerTransport, *, task_status: TaskStatus[None]) -> None: + run_calls.append(None) + task_status.started() + raise TestException("Simulated crash") sent_messages: list[Message] = [] @@ -331,7 +341,7 @@ async def mock_receive(): return {"type": "http.request", "body": b"", "more_body": False} # Trigger session creation - with patch("mcp.server.streamable_http_manager.serve_loop", mock_serve): + with patch.object(StreamableHTTPServerTransport, "run", mock_run): await manager.handle_request(scope, mock_receive, mock_send) session_id = None @@ -346,7 +356,7 @@ async def mock_receive(): assert session_id is not None, "Session ID not found in response headers" - mock_serve.assert_called_once() + assert len(run_calls) == 1 # Give other tasks a chance to run to ensure the finally block executes await anyio.sleep(0.01) @@ -412,8 +422,9 @@ async def mock_receive(): # The key assertion - transport should be terminated assert transport._terminated, "Transport should be terminated after stateless request" - # Verify internal state is cleaned up - assert len(transport._request_streams) == 0, "Transport should have no active request streams" + # Verify internal state is cleaned up: no request streams left open. + assert not transport._streams, "Transport should have no active request streams" + assert not transport._standalone.attached, "Transport should have no standalone stream attached" @pytest.mark.anyio diff --git a/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py deleted file mode 100644 index 3086dca990..0000000000 --- a/tests/server/test_streamable_http_router.py +++ /dev/null @@ -1,116 +0,0 @@ -"""Regression coverage for the StreamableHTTP per-session response router.""" - -import anyio -import pytest -from mcp_types import JSONRPCMessage, JSONRPCResponse -from starlette.types import Message, Scope - -from mcp.server.streamable_http import ( - REQUEST_STREAM_BUFFER_SIZE, - EventCallback, - EventId, - EventMessage, - EventStore, - StreamableHTTPServerTransport, - StreamId, -) -from mcp.shared.message import SessionMessage - - -class _PrimingFailingStore(EventStore): - async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: - raise RuntimeError("backend unavailable") - - async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: - raise NotImplementedError - - -@pytest.mark.anyio -async def test_router_unconsumed_request_stream_does_not_block_siblings() -> None: - """A response whose `sse_writer` is not yet receiving must not park the router (#1764). - - Drives the routing layer directly (the production race does not reproduce - on loopback), so this pins the router semantics, not the call sites. - """ - transport = StreamableHTTPServerTransport(mcp_session_id="sid", is_json_response_enabled=False) - streams = transport._request_streams - async with transport.connect() as (_read_stream, write_stream): - # Model two concurrent POSTs at the point _handle_post_request has - # registered the per-request stream but A's sse_writer has not yet - # reached its first receive(). - streams["A"] = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) - streams["B"] = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) - a_send, a_recv = streams["A"] - b_reader = streams["B"][1] - b_received = anyio.Event() - - async def consume_b() -> None: - async with b_reader: - await b_reader.receive() - b_received.set() - - async def server_writes() -> None: - await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="A", result={}))) - await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="B", result={}))) - - async with anyio.create_task_group() as tg: - tg.start_soon(consume_b) - tg.start_soon(server_writes) - with anyio.fail_after(5): - await b_received.wait() - # A's response was buffered for its (late) consumer, not dropped. - assert a_send.statistics().current_buffer_used == 1 - await a_recv.aclose() - await a_send.aclose() - - -@pytest.mark.anyio -async def test_priming_store_failure_leaves_no_per_request_state() -> None: - """`EventStore.store_event` raising on the priming row must not leak per-request entries.""" - transport = StreamableHTTPServerTransport( - mcp_session_id=None, - is_json_response_enabled=False, - event_store=_PrimingFailingStore(), - ) - - body = b'{"jsonrpc":"2.0","id":"req-1","method":"tools/list","params":{}}' - scope: Scope = { - "type": "http", - "method": "POST", - "path": "/", - "query_string": b"", - "headers": [ - (b"accept", b"application/json, text/event-stream"), - (b"content-type", b"application/json"), - (b"mcp-protocol-version", b"2025-11-25"), - ], - } - body_sent = False - - async def receive() -> Message: - nonlocal body_sent - if not body_sent: - body_sent = True - return {"type": "http.request", "body": body, "more_body": False} - raise NotImplementedError - - sent: list[Message] = [] - - async def asgi_send(message: Message) -> None: - sent.append(message) - - async with transport.connect() as (read_stream, _write_stream): - async with anyio.create_task_group() as tg: - tg.start_soon(transport.handle_request, scope, receive, asgi_send) - with anyio.fail_after(5): - forwarded = await read_stream.receive() - assert isinstance(forwarded, Exception) - # handle_request has returned; connect()'s finally (which clears - # _request_streams unconditionally) has not yet run. - assert transport._request_streams == {} - assert transport._sse_stream_writers == {} - - assert sent[0]["type"] == "http.response.start" - assert sent[0]["status"] == 500 - body = b"".join(m.get("body", b"") for m in sent if m["type"] == "http.response.body") - assert b"backend unavailable" not in body diff --git a/tests/server/test_streamable_http_transport.py b/tests/server/test_streamable_http_transport.py new file mode 100644 index 0000000000..a129988314 --- /dev/null +++ b/tests/server/test_streamable_http_transport.py @@ -0,0 +1,884 @@ +"""Behaviour of the StreamableHTTP server transport's per-request dispatch. + +Each POSTed request is served by its own response channel and a session-scoped +correlator; these tests pin the parts of that lifecycle a real client can hit +that the transport-agnostic interaction matrix does not reach. +""" + +from typing import Any +from unittest.mock import MagicMock + +import anyio +import anyio.lowlevel +import pytest +from httpx2 import EventSource +from mcp_types import ( + CONNECTION_CLOSED, + INVALID_REQUEST, + CallToolRequestParams, + CallToolResult, + ElicitRequest, + ElicitRequestFormParams, + ElicitResult, + JSONRPCError, + JSONRPCMessage, + JSONRPCNotification, + JSONRPCRequest, + JSONRPCResponse, + TextContent, +) +from starlette.requests import Request +from starlette.types import Message, Scope + +from mcp.server import Server, ServerRequestContext +from mcp.server.context import CallNext, HandlerResult +from mcp.server.streamable_http import ( + EventCallback, + EventId, + EventStore, + StreamableHTTPServerTransport, + StreamId, + _HTTPRequestDispatchContext, # pyright: ignore[reportPrivateUsage] + _MessageChannel, # pyright: ignore[reportPrivateUsage] +) +from mcp.shared._correlation import RequestCorrelator +from mcp.shared.exceptions import MCPError, NoBackChannelError +from mcp.shared.message import ServerMessageMetadata +from mcp.shared.transport_context import TransportContext +from tests.interaction._connect import ( + base_headers, + connect_over_streamable_http, + initialize_body, + initialize_via_http, + mounted_app, + parse_sse_messages, +) +from tests.interaction.transports._event_store import SequencedEventStore + +pytestmark = pytest.mark.anyio + + +class _PrimingFailingStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + raise RuntimeError("backend unavailable") + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise NotImplementedError + + +class _StreamFailingStore(SequencedEventStore): + """A store that breaks for every message on request ``42``'s stream (its priming row aside).""" + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + if stream_id.endswith(":request:42") and message is not None: + raise RuntimeError("backend fell over") + return await super().store_event(stream_id, message) + + +def _tools_call(request_id: int, name: str, arguments: dict[str, object]) -> str: + return JSONRPCRequest( + jsonrpc="2.0", id=request_id, method="tools/call", params={"name": name, "arguments": arguments} + ).model_dump_json(by_alias=True, exclude_none=True) + + +async def test_priming_store_failure_returns_500_without_leaking_per_request_state() -> None: + """`EventStore.store_event` raising on the priming row yields a 500 with no leaked state or backend text. + + The priming row is minted before any per-request state exists, so a failing + store leaves nothing to clean up and its exception text never reaches the wire. + """ + transport = StreamableHTTPServerTransport( + mcp_session_id=None, + is_json_response_enabled=False, + event_store=_PrimingFailingStore(), + app=Server("priming-failure"), + lifespan_state={}, + ) + + body = b'{"jsonrpc":"2.0","id":"req-1","method":"tools/list","params":{}}' + scope: Scope = { + "type": "http", + "method": "POST", + "path": "/", + "query_string": b"", + "headers": [ + (b"accept", b"application/json, text/event-stream"), + (b"content-type", b"application/json"), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + body_sent = False + + async def receive() -> Message: + nonlocal body_sent + if not body_sent: + body_sent = True + return {"type": "http.request", "body": body, "more_body": False} + raise NotImplementedError + + sent: list[Message] = [] + + async def asgi_send(message: Message) -> None: + sent.append(message) + + with anyio.fail_after(5): + await transport.handle_request(scope, receive, asgi_send) + + assert transport._streams == {} # pyright: ignore[reportPrivateUsage] + assert sent[0]["type"] == "http.response.start" + assert sent[0]["status"] == 500 + payload = b"".join(m.get("body", b"") for m in sent if m["type"] == "http.response.body") + assert b"backend unavailable" not in payload + + +async def test_terminating_a_session_ends_its_in_flight_request_streams_and_cancels_the_handlers() -> None: + """DELETE while a call is running closes that call's SSE stream and cancels its handler.""" + started = anyio.Event() + cancelled = anyio.Event() + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + started.set() + try: + await anyio.sleep_forever() + finally: + cancelled.set() + raise NotImplementedError # unreachable: the handler is cancelled while sleeping + + server = Server("terminating", on_call_tool=call_tool) + + async with mounted_app(server) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( + "POST", "/mcp", content=_tools_call(1, "wait", {}), headers=base_headers(session_id=session_id) + ) as response: + assert response.status_code == 200 + await started.wait() + delete = await http.delete("/mcp", headers=base_headers(session_id=session_id)) + assert delete.status_code == 200 + # Termination closes the request's stream, so the read ends here. + events = [event async for event in EventSource(response)] + await cancelled.wait() + follow_up = await http.post( + "/mcp", + json={"jsonrpc": "2.0", "id": 2, "method": "ping"}, + headers=base_headers(session_id=session_id), + ) + + assert all(not isinstance(message, JSONRPCResponse) for message in parse_sse_messages(events)) + assert follow_up.status_code == 404 + + +async def test_a_posted_progress_notification_reaches_the_servers_pending_request() -> None: + """A client POSTs notifications/progress for a request the server sent it; the server's callback receives it. + + The elicitation request rides the tool call's own SSE stream (related to it); the progress + notification and the answer arrive as separate POSTs and are correlated back to the pending + request by the token / id the server minted. + """ + reports: list[tuple[float, float | None, str | None]] = [] + + async def on_progress(progress: float, total: float | None, message: str | None) -> None: + reports.append((progress, total, message)) + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + result = await ctx.session.send_request( + ElicitRequest( + params=ElicitRequestFormParams(message="ok?", requested_schema={"type": "object", "properties": {}}) + ), + ElicitResult, + metadata=ServerMessageMetadata(related_request_id=ctx.request_id), + progress_callback=on_progress, + ) + return CallToolResult(content=[TextContent(text=result.action)]) + + server = Server("progressive", on_call_tool=call_tool) + + async with mounted_app(server) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( # pragma: no branch + "POST", "/mcp", content=_tools_call(1, "ask", {}), headers=base_headers(session_id=session_id) + ) as response: + assert response.status_code == 200 + events = aiter(EventSource(response)) + elicit_event = await anext(events) + elicit = JSONRPCRequest.model_validate_json(elicit_event.data) + assert elicit.method == "elicitation/create" + assert elicit.params is not None + token = elicit.params["_meta"]["progressToken"] + progress = await http.post( + "/mcp", + json={ + "jsonrpc": "2.0", + "method": "notifications/progress", + "params": {"progressToken": token, "progress": 0.5, "total": 1.0, "message": "half"}, + }, + headers=base_headers(session_id=session_id), + ) + assert progress.status_code == 202 + answer = await http.post( + "/mcp", + json={"jsonrpc": "2.0", "id": elicit.id, "result": {"action": "accept", "content": {}}}, + headers=base_headers(session_id=session_id), + ) + assert answer.status_code == 202 + result_event = await anext(events) + result = JSONRPCResponse.model_validate_json(result_event.data) + assert result.result["content"] == [{"type": "text", "text": "accept"}] + assert reports == [(0.5, 1.0, "half")] + + +async def test_an_event_store_failure_costs_that_request_its_resumability_only() -> None: + """A store that raises for a request's stream still lets the answer through live. + + Request 42 cannot be stored, so its result reaches the client unstored (no event id to + resume from) and the store's error text never touches the wire; request 43 on the same + session is stored and served normally. + """ + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + return CallToolResult(content=[TextContent(text=params.name)]) + + server = Server("resilient", on_call_tool=call_tool) + store = _StreamFailingStore() + + async with mounted_app(server, event_store=store, retry_interval=0) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( + "POST", "/mcp", content=_tools_call(42, "first", {}), headers=base_headers(session_id=session_id) + ) as failing: + assert failing.status_code == 200 + failing_events = [event async for event in EventSource(failing)] + async with http.stream( # pragma: no branch + "POST", "/mcp", content=_tools_call(43, "second", {}), headers=base_headers(session_id=session_id) + ) as healthy: + assert healthy.status_code == 200 + healthy_events = [event async for event in EventSource(healthy)] + + # Request 42's answer went out live (the store never took it) and no backend text leaked. + stored_ids = {message.id for _, message in store._events if isinstance(message, JSONRPCResponse)} # pyright: ignore[reportPrivateUsage] + assert 42 not in stored_ids and 43 in stored_ids + (first,) = [message for message in parse_sse_messages(failing_events) if isinstance(message, JSONRPCResponse)] + assert first.id == 42 + assert first.result["content"] == [{"type": "text", "text": "first"}] + assert all("fell over" not in (event.data or "") for event in failing_events) + # Request 43 was stored and delivered as usual. + (second,) = [message for message in parse_sse_messages(healthy_events) if isinstance(message, JSONRPCResponse)] + assert second.id == 43 + assert second.result["content"] == [{"type": "text", "text": "second"}] + + +async def test_concurrent_posts_reusing_a_request_id_each_receive_their_own_response() -> None: + """Two concurrent requests sharing a JSON-RPC id are answered on their own POST streams. + + Each POST owns its response channel, so the second registration does not steal or clobber + the first's stream (the session-level entry only serves close/replay lookup). + """ + slow_started = anyio.Event() + release = anyio.Event() + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + if params.name == "slow": + slow_started.set() + await release.wait() + return CallToolResult(content=[TextContent(text=params.name)]) + + server = Server("dupes", on_call_tool=call_tool) + results: dict[str, JSONRPCResponse] = {} + + async with mounted_app(server) as (http, _): + session_id = await initialize_via_http(http) + + async def post(name: str) -> None: + async with http.stream( + "POST", "/mcp", content=_tools_call(7, name, {}), headers=base_headers(session_id=session_id) + ) as response: + events = [event async for event in EventSource(response)] + (message,) = parse_sse_messages(events) + assert isinstance(message, JSONRPCResponse) + results[name] = message + + with anyio.fail_after(5): + async with anyio.create_task_group() as tg: # pragma: no branch + tg.start_soon(post, "slow") + await slow_started.wait() + await post("fast") + release.set() + + assert results["fast"].result["content"] == [{"type": "text", "text": "fast"}] + assert results["slow"].result["content"] == [{"type": "text", "text": "slow"}] + assert {results["fast"].id, results["slow"].id} == {7} + + +def test_detaching_a_stale_attachment_does_not_evict_the_newer_one() -> None: + """A response that finishes after a Last-Event-ID reconnect re-attached must not knock the newer + attachment off its channel.""" + channel = _MessageChannel("1", None) + stale_reader = channel.attach() + assert stale_reader is not None + channel.detach() # e.g. close_sse_stream() + fresh_reader = channel.attach() # the client's reconnect re-attached + assert fresh_reader is not None + channel.detach(stale_reader) # the stale response's cleanup lands late + assert channel.attached + channel.detach(fresh_reader) + assert not channel.attached + stale_reader.close() + fresh_reader.close() + + +def test_closing_streams_for_unknown_requests_is_a_no_op() -> None: + """`close_sse_stream` / `close_standalone_sse_stream` with nothing open do nothing.""" + transport = StreamableHTTPServerTransport("sid") + transport.close_sse_stream("no-such-request") + transport.close_standalone_sse_stream() + + +async def test_a_transport_not_bound_to_a_server_refuses_to_handle_requests() -> None: + """The transport is created by the session manager; driving one built without a server fails loudly.""" + transport = StreamableHTTPServerTransport("sid") + scope: Scope = {"type": "http", "method": "GET", "path": "/", "query_string": b"", "headers": []} + + async def receive() -> Message: + raise NotImplementedError + + async def send(message: Message) -> None: + raise NotImplementedError + + with pytest.raises(RuntimeError, match="not bound to a server"): + await transport.handle_request(scope, receive, send) + + +async def test_a_closed_request_context_drops_notifications_and_refuses_requests() -> None: + """Once the handler has returned, its context stops accepting output (a background task can't + write onto a finished request's stream).""" + channel = _MessageChannel("1", None) + dctx = _HTTPRequestDispatchContext( + transport=TransportContext(kind="streamable-http", can_send_request=True), + _corr=RequestCorrelator(), + _channel=channel, + _request_id=1, + ) + assert dctx.can_send_request + reader = channel.attach() + assert reader is not None + + dctx.close() + + await dctx.notify("notifications/message", {"level": "info", "data": "too late"}) + with pytest.raises(anyio.WouldBlock): + reader.receive_nowait() # nothing reached the response stream + reader.close() + channel.close() + with pytest.raises(NoBackChannelError): + await dctx.send_raw_request("ping", None) + + +class _SlowFirstStore(EventStore): + """The first `store_event` call is slow, so an unordered second writer could overtake it.""" + + def __init__(self) -> None: + self.count = 0 + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + self.count += 1 + event_id = str(self.count) + if event_id == "1": + for _ in range(5): + await anyio.lowlevel.checkpoint() + return event_id + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise NotImplementedError + + +async def test_concurrent_writes_reach_the_wire_in_event_store_order() -> None: + """Two tasks writing one channel are delivered in the order the event store recorded them. + + A `Last-Event-ID` resume replays in store order, so the wire must never diverge from it. + """ + channel = _MessageChannel("1", _SlowFirstStore()) + reader = channel.attach() + assert reader is not None + first = JSONRPCNotification(jsonrpc="2.0", method="notifications/one") + second = JSONRPCNotification(jsonrpc="2.0", method="notifications/two") + + with anyio.fail_after(5): + async with anyio.create_task_group() as tg: + tg.start_soon(channel.write, first) + await anyio.lowlevel.checkpoint() # let the first writer reach the store + tg.start_soon(channel.write, second) + + delivered = [reader.receive_nowait(), reader.receive_nowait()] + reader.close() + channel.close() + assert [(event.event_id, event.message) for event in delivered] == [("1", first), ("2", second)] + + +class _GatedPrimingStore(SequencedEventStore): + """Parks request ``42``'s priming write until released, so a DELETE can land mid-request.""" + + def __init__(self) -> None: + super().__init__() + self.parked = anyio.Event() + self.release = anyio.Event() + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + if stream_id.endswith(":request:42") and message is None: + self.parked.set() + await self.release.wait() + return await super().store_event(stream_id, message) + + +async def test_a_request_arriving_across_termination_is_refused_not_run() -> None: + """A POST suspended when its session is DELETEd is answered 404 rather than run on the dead session.""" + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + raise NotImplementedError # the handler must never run + + server = Server("terminating", on_call_tool=call_tool) + store = _GatedPrimingStore() + + pending: list[int] = [] + async with mounted_app(server, event_store=store, retry_interval=0) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with anyio.create_task_group() as tg: # pragma: no branch + + async def call() -> None: + response = await http.post( + "/mcp", content=_tools_call(42, "wait", {}), headers=base_headers(session_id=session_id) + ) + pending.append(response.status_code) + + tg.start_soon(call) + await store.parked.wait() + delete = await http.delete("/mcp", headers=base_headers(session_id=session_id)) + assert delete.status_code == 200 + store.release.set() + + assert pending == [404] + + +class _BrokenReplayStore(SequencedEventStore): + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise RuntimeError("replay backend unavailable") + + +async def test_a_failing_replay_ends_that_stream_and_the_session_keeps_working() -> None: + """`replay_events_after` raising costs the reconnecting GET an empty stream, nothing more.""" + server = Server("replay-broken") + + async with mounted_app(server, event_store=_BrokenReplayStore(), retry_interval=0) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( # pragma: no branch + "GET", "/mcp", headers=base_headers(session_id=session_id) | {"last-event-id": "1"} + ) as replay: + assert replay.status_code == 200 + assert [event async for event in EventSource(replay)] == [] + ping = await http.post( + "/mcp", + json={"jsonrpc": "2.0", "id": 2, "method": "ping"}, + headers=base_headers(session_id=session_id), + ) + assert ping.status_code == 200 + + +async def test_json_mode_refuses_a_request_scoped_server_to_client_request() -> None: + """In JSON-response mode an elicitation from a handler fails with a JSON-RPC error, not a hang. + + The POST's single JSON body has no stream to carry the nested request, so the transport + raises `NoBackChannelError` (an `MCPError`) rather than waiting for an answer that could + never be delivered. + """ + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + await ctx.session.send_request( + ElicitRequest( + params=ElicitRequestFormParams(message="ok?", requested_schema={"type": "object", "properties": {}}) + ), + ElicitResult, + metadata=ServerMessageMetadata(related_request_id=ctx.request_id), + ) + raise NotImplementedError # the request must be refused before reaching here + + server = Server("json-mode", on_call_tool=call_tool) + + async with connect_over_streamable_http(server, json_response=True) as client: + with pytest.raises(MCPError) as exc_info, anyio.fail_after(5): + await client.call_tool("ask", {}) + + assert exc_info.value.error.code == INVALID_REQUEST + + +async def test_a_session_bound_transport_drops_notifications_once_its_session_has_ended() -> None: + """A POSTed notification landing after the session task is gone is dropped, not handled.""" + transport = StreamableHTTPServerTransport("sid", app=Server("ended"), lifespan_state={}) + # The session task (`run()`) never started, so the transport has no live session to hand work to. + request = MagicMock(spec=Request) + request.headers = {} + + with anyio.fail_after(5): + await transport._deliver_client_message( # pyright: ignore[reportPrivateUsage] + request, JSONRPCNotification(jsonrpc="2.0", method="notifications/initialized") + ) + + +async def test_requests_wait_for_a_handshake_in_progress_to_commit() -> None: + """A request POSTed while `initialize` is still running is held until the handshake commits. + + The session id ships with the initialize response's headers, ahead of its result, so a + client can send its next request before the server committed the negotiated session state. + The transport orders that request behind the handshake, so it is served against the + initialized session rather than racing the initialization gate. + """ + release_handshake = anyio.Event() + + class _SlowHandshake: + async def __call__(self, ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult: + if ctx.method == "initialize": + await release_handshake.wait() + return await call_next(ctx) + + server = Server("gated") + server.middleware.append(_SlowHandshake()) + listed: list[int] = [] + + async with mounted_app(server) as (http, _): + with anyio.fail_after(5): + async with anyio.create_task_group() as tg: # pragma: no branch + async with http.stream( # pragma: no branch + "POST", "/mcp", json=initialize_body(), headers=base_headers() + ) as init: + # The session id arrives with the headers, before the handshake finishes. + session_id = init.headers["mcp-session-id"] + + async def list_tools() -> None: + response = await http.post( + "/mcp", + json={"jsonrpc": "2.0", "id": 2, "method": "tools/list"}, + headers=base_headers(session_id=session_id), + ) + listed.append(response.status_code) + + tg.start_soon(list_tools) + await anyio.wait_all_tasks_blocked() + assert listed == [] # held behind the still-running handshake + release_handshake.set() + [event async for event in EventSource(init)] + + assert listed == [200] + + +async def test_a_server_request_no_client_can_receive_fails_the_call_instead_of_hanging() -> None: + """A server-to-client request with nothing to carry it fails the handler right away. + + With no GET stream attached and no event store there is nowhere a connection-scoped + request could ever reach the client, so the call is failed `CONNECTION_CLOSED` instead + of parking the handler for an answer that cannot arrive. + """ + outcomes: list[int] = [] + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + try: + # No `related_request_id`: this rides the connection's standalone GET stream. + await ctx.session.send_request( + ElicitRequest( + params=ElicitRequestFormParams(message="ok?", requested_schema={"type": "object", "properties": {}}) + ), + ElicitResult, + ) + except MCPError as exc: + outcomes.append(exc.error.code) + raise + raise NotImplementedError + + server = Server("no-back-channel", on_call_tool=call_tool) + + async with mounted_app(server) as (http, _): # no event store, and no GET stream is ever opened + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( # pragma: no branch + "POST", "/mcp", content=_tools_call(2, "ask", {}), headers=base_headers(session_id=session_id) + ) as response: + assert response.status_code == 200 + events = [event async for event in EventSource(response)] + + assert outcomes == [CONNECTION_CLOSED] + # The tool's failure came back on its own stream as an error frame. + (error,) = [message for message in parse_sse_messages(events) if isinstance(message, JSONRPCError)] + assert error.id == 2 + + +class _StandaloneFlakyStore(SequencedEventStore): + """The store rejects the first message written to the standalone stream, then recovers.""" + + def __init__(self) -> None: + super().__init__() + self.failed_once = False + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + if stream_id.endswith(":_GET_stream") and message is not None and not self.failed_once: + self.failed_once = True + raise RuntimeError("standalone backend hiccup") + return await super().store_event(stream_id, message) + + +async def test_a_store_failure_on_the_standalone_stream_does_not_take_the_stream_down() -> None: + """A store that raises for a standalone notification degrades that message, not the GET stream. + + The failed notification still reaches the connected client live (unstored); the next one + is stored and delivered too - the stream stays alive across the store's hiccup. + """ + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + await ctx.session.send_resource_updated("file:///one") + await ctx.session.send_resource_updated("file:///two") + return CallToolResult(content=[]) + + server = Server("standalone-flaky", on_call_tool=call_tool) + store = _StandaloneFlakyStore() + + async with mounted_app(server, event_store=store, retry_interval=0) as (http, _): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( # pragma: no branch + "GET", "/mcp", headers=base_headers(session_id=session_id) + ) as sse: + assert sse.status_code == 200 + events = aiter(EventSource(sse)) + called = await http.post( + "/mcp", content=_tools_call(2, "log", {}), headers=base_headers(session_id=session_id) + ) + assert called.status_code == 200 + updated = [ + JSONRPCNotification.model_validate_json((await anext(events)).data or "{}") for _ in range(2) + ] + + assert [n.params and n.params["uri"] for n in updated] == ["file:///one", "file:///two"] + assert store.failed_once + + +async def test_a_shared_event_store_keeps_sessions_replay_apart() -> None: + """Two sessions on one store never see each other's frames on a `Last-Event-ID` resume. + + Both sessions run the same request ids, so a store keyed on the bare request id would + interleave them; session-scoped stream ids keep session B's replay to its own messages. + """ + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + assert params.arguments is not None + return CallToolResult(content=[TextContent(text=str(params.arguments["owner"]))]) + + server = Server("shared-store", on_call_tool=call_tool) + store = SequencedEventStore() + + async with mounted_app(server, event_store=store, retry_interval=0) as (http, _): + with anyio.fail_after(5): + session_a = await initialize_via_http(http) + session_b = await initialize_via_http(http) + events_by_session: dict[str, list[Any]] = {} + for session_id, owner in ((session_a, "A"), (session_b, "B")): + async with http.stream( + "POST", + "/mcp", + content=_tools_call(3, "who", {"owner": owner}), + headers=base_headers(session_id=session_id), + ) as response: + events_by_session[session_id] = [event async for event in EventSource(response)] + # Resume session B's request-3 stream from its priming event: only B's frames come back. + priming_a = [event.id for event in events_by_session[session_a] if event.id][0] + priming_b = [event.id for event in events_by_session[session_b] if event.id][0] + async with http.stream( # pragma: no branch + "GET", + "/mcp", + headers=base_headers(session_id=session_b) | {"last-event-id": priming_b}, + ) as replay: + replayed = [event async for event in EventSource(replay)] + # ... while an event id belonging to session A yields B nothing. + async with http.stream( # pragma: no branch + "GET", + "/mcp", + headers=base_headers(session_id=session_b) | {"last-event-id": priming_a}, + ) as poached: + foreign = [event async for event in EventSource(poached)] + + payloads = [message for message in parse_sse_messages(replayed) if isinstance(message, JSONRPCResponse)] + assert [message.result["content"][0]["text"] for message in payloads] == ["B"] + # Presenting session A's event id from session B replays nothing at all. + assert [message for message in parse_sse_messages(foreign) if isinstance(message, JSONRPCResponse)] == [] + + +async def test_a_request_id_shaped_like_the_get_marker_keeps_its_own_stream() -> None: + """A request whose id is the string `_GET_stream` is served on its own stream. + + Its response never lands on the standalone GET stream, and the standalone stream stays + attached across it. + """ + server = Server("marker-id") + + async with mounted_app(server) as (http, manager): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + async with http.stream( # pragma: no branch + "GET", "/mcp", headers=base_headers(session_id=session_id) + ) as standalone: + assert standalone.status_code == 200 + async with http.stream( # pragma: no branch + "POST", + "/mcp", + json={"jsonrpc": "2.0", "id": "_GET_stream", "method": "ping"}, + headers=base_headers(session_id=session_id), + ) as pinged: + assert pinged.status_code == 200 + (response,) = parse_sse_messages([event async for event in EventSource(pinged)]) + assert isinstance(response, JSONRPCResponse) and response.id == "_GET_stream" + # The standalone stream is still attached; the ping never touched it. + transport = manager._server_instances[session_id] # pyright: ignore[reportPrivateUsage] + assert transport._standalone.attached # pyright: ignore[reportPrivateUsage] + + +async def test_a_json_mode_request_terminated_underneath_it_is_answered_404() -> None: + """DELETE while a JSON-mode request runs leaves it no answer, so the POST gets the terminated 404.""" + started = anyio.Event() + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + started.set() + await anyio.sleep_forever() + raise NotImplementedError + + server = Server("json-terminate", on_call_tool=call_tool) + statuses: list[int] = [] + + async with mounted_app(server, json_response=True) as (http, _): + with anyio.fail_after(5): + # JSON mode answers `initialize` with a plain JSON body. + initialized = await http.post("/mcp", json=initialize_body(), headers=base_headers()) + assert initialized.status_code == 200 + session_id = initialized.headers["mcp-session-id"] + async with anyio.create_task_group() as tg: # pragma: no branch + + async def call() -> None: + response = await http.post( + "/mcp", content=_tools_call(2, "wait", {}), headers=base_headers(session_id=session_id) + ) + statuses.append(response.status_code) + + tg.start_soon(call) + await started.wait() + delete = await http.delete("/mcp", headers=base_headers(session_id=session_id)) + assert delete.status_code == 200 + + assert statuses == [404] + + +class _GatedReplayStore(SequencedEventStore): + """Parks `replay_events_after` until released, so a DELETE can land mid-replay.""" + + def __init__(self) -> None: + super().__init__() + self.parked = anyio.Event() + self.release = anyio.Event() + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + self.parked.set() + await self.release.wait() + return await super().replay_events_after(last_event_id, send_callback) + + +async def test_a_replay_across_termination_ends_instead_of_tailing_a_dead_stream() -> None: + """DELETE while a standalone-stream replay reads the store ends that response cleanly. + + The resumed GET finds its stream gone once the store read returns, so it closes with no + live tail rather than attaching to a terminated channel. + """ + + async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: + await ctx.session.send_resource_updated("file:///seed") + return CallToolResult(content=[]) + + server = Server("replay-terminate", on_call_tool=call_tool) + store = _GatedReplayStore() + replayed: list[list[Any]] = [] + + async with mounted_app(server, event_store=store, retry_interval=0) as (http, manager): + session_id = await initialize_via_http(http) + with anyio.fail_after(5): + # Seed the standalone stream with one stored event, and note its id. + async with http.stream("GET", "/mcp", headers=base_headers(session_id=session_id)) as sse: + events = aiter(EventSource(sse)) + seeded = await http.post( + "/mcp", content=_tools_call(2, "seed", {}), headers=base_headers(session_id=session_id) + ) + assert seeded.status_code == 200 + last_event_id = (await anext(events)).id + assert last_event_id is not None + transport = manager._server_instances[session_id] # pyright: ignore[reportPrivateUsage] + while transport._standalone.attached: # pyright: ignore[reportPrivateUsage] + await anyio.wait_all_tasks_blocked() # let the closed GET detach + + async with anyio.create_task_group() as tg: # pragma: no branch + + async def replay() -> None: + async with http.stream( + "GET", + "/mcp", + headers=base_headers(session_id=session_id) | {"last-event-id": last_event_id}, + ) as response: + replayed.append([event async for event in EventSource(response)]) + + tg.start_soon(replay) + await store.parked.wait() + delete = await http.delete("/mcp", headers=base_headers(session_id=session_id)) + assert delete.status_code == 200 + store.release.set() + + (events,) = replayed + assert [event for event in events if event.data] == [] + + +async def test_overlapping_handshakes_each_release_their_own_gate() -> None: + """A second `initialize` POSTed while the first runs installs its own gate; both are answered. + + Whichever handshake finishes clears only its own gate, so neither leaves a stale gate + holding later requests forever. + """ + release_handshake = anyio.Event() + + class _SlowHandshake: + async def __call__(self, ctx: ServerRequestContext, call_next: CallNext) -> HandlerResult: + if ctx.method == "initialize" and ctx.request_id != 1: # let the session's first one through + await release_handshake.wait() + return await call_next(ctx) + + server = Server("regated") + server.middleware.append(_SlowHandshake()) + statuses: list[int] = [] + + async def handshake(http: Any, session_id: str, request_id: int) -> None: + response = await http.post( + "/mcp", json=initialize_body(request_id), headers=base_headers(session_id=session_id) + ) + statuses.append(response.status_code) + + async with mounted_app(server, json_response=True) as (http, _): + with anyio.fail_after(5): + first = await http.post("/mcp", json=initialize_body(), headers=base_headers()) + session_id = first.headers["mcp-session-id"] + async with anyio.create_task_group() as tg: # pragma: no branch + tg.start_soon(handshake, http, session_id, 2) + await anyio.wait_all_tasks_blocked() + tg.start_soon(handshake, http, session_id, 3) + await anyio.wait_all_tasks_blocked() + release_handshake.set() + listed = await http.post( + "/mcp", + json={"jsonrpc": "2.0", "id": 4, "method": "tools/list"}, + headers=base_headers(session_id=session_id), + ) + + assert statuses == [200, 200] + assert listed.status_code == 200 diff --git a/tests/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index aeef25a278..8ccbc40683 100644 --- a/tests/shared/test_streamable_http.py +++ b/tests/shared/test_streamable_http.py @@ -7,6 +7,7 @@ from __future__ import annotations as _annotations import json +import logging import time from collections.abc import AsyncIterator from contextlib import asynccontextmanager @@ -19,7 +20,6 @@ import httpx2 import mcp_types as types import pytest -from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from httpx2 import ServerSentEvent from mcp_types import ( DEFAULT_NEGOTIATED_VERSION, @@ -40,7 +40,6 @@ from starlette.applications import Starlette from starlette.requests import Request from starlette.routing import Mount -from starlette.types import Message, Scope from mcp import MCPError from mcp.client import ClientRequestContext, IncomingMessage @@ -48,7 +47,6 @@ from mcp.client.streamable_http import StreamableHTTPTransport, streamable_http_client from mcp.server import Server, ServerRequestContext from mcp.server.streamable_http import ( - GET_STREAM_KEY, MCP_PROTOCOL_VERSION_HEADER, MCP_SESSION_ID_HEADER, SESSION_ID_PATTERN, @@ -1639,27 +1637,24 @@ async def test_close_sse_stream_callback_not_provided_for_old_protocol_version() event_store=SimpleEventStore(), ) - # Create a mock message and request - mock_message = JSONRPCRequest(jsonrpc="2.0", id="test-1", method="tools/list") + # Create a mock request mock_request = MagicMock() - # Call _create_session_message with OLD protocol version - session_msg = transport._create_session_message(mock_message, mock_request, "test-request-id", "2025-06-18") + # Build the per-request metadata with OLD protocol version + metadata = transport._build_message_metadata(mock_request, "test-request-id", "2025-06-18") # Callbacks should NOT be provided for old protocol version - assert session_msg.metadata is not None - assert isinstance(session_msg.metadata, ServerMessageMetadata) - assert session_msg.metadata.close_sse_stream is None - assert session_msg.metadata.close_standalone_sse_stream is None + assert isinstance(metadata, ServerMessageMetadata) + assert metadata.close_sse_stream is None + assert metadata.close_standalone_sse_stream is None # Now test with NEW protocol version - should provide callbacks - session_msg_new = transport._create_session_message(mock_message, mock_request, "test-request-id-2", "2025-11-25") + metadata_new = transport._build_message_metadata(mock_request, "test-request-id-2", "2025-11-25") # Callbacks SHOULD be provided for new protocol version - assert session_msg_new.metadata is not None - assert isinstance(session_msg_new.metadata, ServerMessageMetadata) - assert session_msg_new.metadata.close_sse_stream is not None - assert session_msg_new.metadata.close_standalone_sse_stream is not None + assert isinstance(metadata_new, ServerMessageMetadata) + assert metadata_new.close_sse_stream is not None + assert metadata_new.close_standalone_sse_stream is not None @pytest.mark.anyio @@ -1670,15 +1665,13 @@ async def test_close_sse_stream_callback_not_provided_for_unknown_protocol_versi event_store=SimpleEventStore(), ) - mock_message = JSONRPCRequest(jsonrpc="2.0", id="test-1", method="tools/list") mock_request = MagicMock() - session_msg = transport._create_session_message(mock_message, mock_request, "test-request-id", "zzz") + metadata = transport._build_message_metadata(mock_request, "test-request-id", "zzz") - assert session_msg.metadata is not None - assert isinstance(session_msg.metadata, ServerMessageMetadata) - assert session_msg.metadata.close_sse_stream is None - assert session_msg.metadata.close_standalone_sse_stream is None + assert isinstance(metadata, ServerMessageMetadata) + assert metadata.close_sse_stream is None + assert metadata.close_standalone_sse_stream is None @pytest.mark.anyio @@ -2184,83 +2177,5 @@ async def message_handler(message: IncomingMessage) -> None: await notified.wait() # Tear the standalone stream down while the writer is parked on it. (transport,) = session_manager._server_instances.values() # pyright: ignore[reportPrivateUsage] - await transport._clean_up_memory_streams(GET_STREAM_KEY) # pyright: ignore[reportPrivateUsage] - assert "Error in standalone SSE writer" not in caplog.text - - -@pytest.mark.anyio -async def test_standalone_stream_teardown_between_dequeues_is_not_an_error( - caplog: pytest.LogCaptureFixture, -) -> None: - """Teardown landing while the standalone writer is between dequeues logs no error. - - SDK-defined: after teardown the writer's next dequeue hits its own closed stream — expected - disconnect noise. The public surface cannot force this window (the in-process client consumes - SSE without backpressure), so the test drives the transport's ASGI entry point with a gated `send`. - """ - transport = StreamableHTTPServerTransport( - mcp_session_id=None, - security_settings=TransportSecuritySettings(enable_dns_rebinding_protection=False), - ) - # The GET handler only checks that a read-stream writer exists; it is never written to. - read_stream_writer, read_stream = create_context_streams[SessionMessage | Exception](0) - transport._read_stream_writer = read_stream_writer # pyright: ignore[reportPrivateUsage] - - stream_registered = anyio.Event() - - class SignalingStreams( - dict[types.RequestId, tuple[MemoryObjectSendStream[EventMessage], MemoryObjectReceiveStream[EventMessage]]] - ): - # Only the GET handler inserts here, so any insert is the standalone stream registration. - def __setitem__( - self, - key: types.RequestId, - value: tuple[MemoryObjectSendStream[EventMessage], MemoryObjectReceiveStream[EventMessage]], - ) -> None: - super().__setitem__(key, value) - stream_registered.set() - - transport._request_streams = SignalingStreams() # pyright: ignore[reportPrivateUsage] - - gate = anyio.Event() - sent: list[Message] = [] - - async def asgi_send(message: Message) -> None: - sent.append(message) - await gate.wait() - - # Never delivers anything, parking the response's disconnect listener. - disconnect_send, disconnect_receive = anyio.create_memory_object_stream[Message](0) - - async def asgi_receive() -> Message: - return await disconnect_receive.receive() - - scope: Scope = { - "type": "http", - "method": "GET", - "path": "/mcp", - "query_string": b"", - "headers": [(b"accept", b"text/event-stream")], - } - notification = types.JSONRPCNotification(jsonrpc="2.0", method="notifications/initialized") - - async with read_stream_writer, read_stream, disconnect_send, disconnect_receive: - with anyio.fail_after(5): - async with anyio.create_task_group() as tg: # pragma: no branch - tg.start_soon(transport.handle_request, scope, asgi_receive, asgi_send) - await stream_registered.wait() - standalone_send = transport._request_streams[GET_STREAM_KEY][0] # pyright: ignore[reportPrivateUsage] - # Zero-buffer rendezvous: once send() returns, the writer has dequeued the event - # and is blocked forwarding it past the closed gate — the between-dequeues window. - await standalone_send.send(EventMessage(notification)) - await transport._clean_up_memory_streams(GET_STREAM_KEY) # pyright: ignore[reportPrivateUsage] - # Unblock the response; the writer's next dequeue hits its closed stream. - gate.set() - - assert sent[0]["type"] == "http.response.start" - assert sent[0]["status"] == 200 - body_chunks = [message for message in sent if message["type"] == "http.response.body"] - assert b"notifications/initialized" in body_chunks[0]["body"] - assert body_chunks[-1] == {"type": "http.response.body", "body": b"", "more_body": False} - assert "Error in standalone SSE writer" not in caplog.text - assert "Error in standalone SSE response" not in caplog.text + transport.close_standalone_sse_stream() + assert [r for r in caplog.records if r.name == "mcp.server.streamable_http" and r.levelno >= logging.ERROR] == []