From 920f8a8a335f7d07abf01486e12749997cf65f52 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 19:24:51 +0000 Subject: [PATCH 1/7] Extract RequestCorrelator from JSONRPCDispatcher Moves outbound request correlation, the inbound in-flight table, and the exception-to-wire policy into shared/_correlation.py so a transport without a stream pair can share the semantics. JSONRPCDispatcher delegates; behaviour unchanged. --- src/mcp/shared/_correlation.py | 477 +++++++++++++++++++++++++++ src/mcp/shared/jsonrpc_dispatcher.py | 367 ++++----------------- 2 files changed, 536 insertions(+), 308 deletions(-) create mode 100644 src/mcp/shared/_correlation.py diff --git a/src/mcp/shared/_correlation.py b/src/mcp/shared/_correlation.py new file mode 100644 index 0000000000..008d740584 --- /dev/null +++ b/src/mcp/shared/_correlation.py @@ -0,0 +1,477 @@ +"""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]], + raise_handler_exceptions: bool = False, + ) -> None: + """Run one registered inbound request and write exactly 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). 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 + 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] + # 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) + if scope.cancelled_caught: + # anyio absorbs the scope's own cancel at __exit__, and + # `cancelled_caught` (unlike `cancel_called`) guarantees the + # result write above did not happen - no double response. + # TODO(L38): spec says SHOULD NOT respond after cancel; + # the existing server always has, so match that for now. + answer_write_started = True + await write_error(ErrorData(code=0, message="Request cancelled")) + 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). The shielded helper + # is needed because bare awaits re-raise here. + if not answer_write_started: + 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 not None: + await write_error(error) + else: + 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. + await write_error(ErrorData(code=0, message=str(e))) + if raise_handler_exceptions: + raise + # 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 42798fdc54..8ff1cd0647 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,9 @@ 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 +30,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,9 +48,7 @@ OnNotify, OnNotifyIntercept, OnRequest, - ProgressFnT, as_request_id, - coerce_request_id, run_notify_intercept, ) from mcp.shared.exceptions import MCPError, NoBackChannelError @@ -69,35 +69,14 @@ 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"] """How `notifications/cancelled` is applied: `"interrupt"` (default) cancels the handler's scope; `"signal"` only sets `ctx.cancel_requested`.""" - -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: @@ -114,23 +93,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.""" @@ -186,20 +148,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: @@ -293,13 +243,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, @@ -321,117 +272,19 @@ async def send_raw_request( RuntimeError: Called before `run()`. """ # Post-close sends get the same CONNECTION_CLOSED contract as in-flight waiters. - if self._closed: + if self._corr.closed: raise MCPError(code=CONNECTION_CLOSED, message="Connection closed") if not self._running: 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, @@ -447,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 @@ -498,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( @@ -574,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. @@ -604,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: @@ -647,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, @@ -675,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, @@ -696,65 +513,18 @@ 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. + The exception-to-wire policy lives in `RequestCorrelator.serve_inbound`; + this only binds the wire writes for a stream-pair transport. """ - answer_write_started = False - 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] - # 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) - if scope.cancelled_caught: - # anyio absorbs the scope's own cancel at __exit__, and - # `cancelled_caught` (unlike `cancel_called`) guarantees the - # result write above did not happen - no double response. - # TODO(L38): spec says SHOULD NOT respond after cancel; - # the existing server always has, so match that for now. - answer_write_started = True - await self._write_error(req.id, ErrorData(code=0, message="Request cancelled")) - 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). The shielded helper - # is needed because bare awaits re-raise here. - if not answer_write_started: - 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 not None: - await self._write_error(req.id, error) - else: - 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. - await self._write_error(req.id, ErrorData(code=0, message=str(e))) - if self._raise_handler_exceptions: - raise - # 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), + 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)) @@ -771,25 +541,6 @@ async def _write_error(self, request_id: RequestId, error: ErrorData) -> None: except (anyio.BrokenResourceError, anyio.ClosedResourceError): logger.debug("dropped error for %r: write stream closed", request_id) - 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) - 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 # the request's own SSE stream instead of a possibly-absent GET stream. From ef159e792b6d22bbe9b803f8a535828781fb8be8 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 19:44:39 +0000 Subject: [PATCH 2/7] Triage transport unit tests around the removed router internals The per-session message router and its stream fan-out are gone, so the tests pinning that mechanism are rewritten against the behaviour they guarded (priming-store failure returns 500 with no leaked state; standalone stream teardown via close_standalone_sse_stream logs no error; the manager evicts sessions whose task exits or crashes; stateless requests leave no channels behind) or deleted where the guarded race is now unrepresentable (#1764 router head-of-line blocking, the standalone writer between-dequeues window). The close_sse_stream protocol-version gating tests move to the renamed metadata builder. --- src/mcp/server/streamable_http.py | 1060 ++++++++++------- src/mcp/server/streamable_http_manager.py | 139 +-- tests/server/test_streamable_http_manager.py | 35 +- tests/server/test_streamable_http_router.py | 116 -- .../server/test_streamable_http_transport.py | 74 ++ tests/shared/test_streamable_http.py | 117 +- 6 files changed, 804 insertions(+), 737 deletions(-) delete mode 100644 tests/server/test_streamable_http_router.py create mode 100644 tests/server/test_streamable_http_transport.py diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index d316345c7e..c37ea2a2a6 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -1,22 +1,35 @@ """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, 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 import anyio +import anyio.abc import pydantic_core from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from mcp_types import ( @@ -28,6 +41,7 @@ ErrorData, JSONRPCError, JSONRPCMessage, + JSONRPCNotification, JSONRPCRequest, JSONRPCResponse, RequestId, @@ -40,11 +54,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,15 +82,17 @@ # 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 # Session ID validation pattern (visible ASCII characters ranging from 0x21 to 0x7E) # Pattern ensures entire string contains only valid characters by using ^ and $ anchors SESSION_ID_PATTERN = re.compile(r"^[\x21-\x7E]+$") +# Streamable HTTP transport kind for `TransportContext.kind`. +STREAMABLE_HTTP_KIND = "streamable-http" + # Type aliases StreamId = str EventId = str @@ -139,18 +163,231 @@ async def replay_events_after( 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. With no + response attached the message is only stored (or, without a store, + dropped with a debug log) - the client can reconnect and replay. + """ + + 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._closed = False + 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 closed by termination).""" + + @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) -> None: + """Store-then-forward one outbound message. Never raises for a dropped connection.""" + if self._closed: + logger.debug("dropped message on closed stream %s", self.stream_id) + return + # 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: + event_id = await self._event_store.store_event(self.stream_id, message) + 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.""" + ) + return + try: + await writer.send(EventMessage(message, event_id)) + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + # The SSE response went away between the attach check and the send. + self._detach(writer) + + def attach(self) -> MemoryObjectReceiveStream[EventMessage]: + """Attach a fresh SSE response and return the reader it drains. + + 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. + """ + writer, reader = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) + if self._writer is not None: + self._writer.close() + self._writer = writer + 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 one attachment: a stale response ending + must not knock a newer (resumed) attachment off the channel. + """ + writer = self._writer + if writer is None: + return + if reader is not None and not self._is_writer_of(writer, reader): + return + self._detach(writer) + + def _is_writer_of( + self, writer: MemoryObjectSendStream[EventMessage], reader: MemoryObjectReceiveStream[EventMessage] + ) -> bool: + # A memory-object stream pair shares one state object. + return getattr(writer, "_state", None) is getattr(reader, "_state", None) + + def _detach(self, writer: MemoryObjectSendStream[EventMessage]) -> None: + if self._writer is writer: + self._writer = None + writer.close() + + def close(self) -> None: + """End the channel outright (session termination): detach and refuse further writes.""" + self._closed = True + if self._writer is not None: + self._detach(self._writer) + self.finished.set() + + def finish(self) -> None: + """Mark the request as finished even if no terminal frame was recorded.""" + self.finished.set() + + +@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 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=channel.write, + 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__( @@ -160,6 +397,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. @@ -176,6 +416,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. @@ -188,17 +432,32 @@ 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() + # The standalone GET stream: server-initiated messages related to no request. + self._standalone = _MessageChannel(GET_STREAM_KEY, event_store) + # In-flight request streams, keyed by stream id (str of the request id), + # so `close_sse_stream()` and `Last-Event-ID` replay can find them. + self._channels: dict[StreamId, _MessageChannel] = {} + # 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 + 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: @@ -223,15 +482,10 @@ 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() + stream_id = str(request_id) + channel = self._standalone if stream_id == GET_STREAM_KEY else self._channels.get(stream_id) + if channel is not None: + channel.detach() def close_standalone_sse_stream(self) -> None: """Close the standalone GET SSE stream, triggering client reconnection. @@ -248,22 +502,44 @@ 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() - def _create_session_message( - self, - message: JSONRPCMessage, - request: Request, - request_id: RequestId, - protocol_version: str, - ) -> SessionMessage: - """Create a session message with metadata including close_sse_stream callback. + 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. + + 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. + """ + try: + async with anyio.create_task_group() as tg: + self._task_group = tg + task_status.started() + await self._closed_event.wait() + tg.cancel_scope.cancel() + finally: + self._task_group = None + # End every still-open response stream and wake anything awaiting a + # client answer; runs on termination and on manager shutdown alike. + for channel in list(self._channels.values()): + channel.close() + self._channels.clear() + self._standalone.close() + self._corr.close() + if self._connection is not None: + await aclose_shielded(self._connection) + + def _build_message_metadata( + self, request: Request, request_id: RequestId, protocol_version: str + ) -> 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. """ - # Only provide close callbacks when client supports resumability if self._event_store and is_version_at_least(protocol_version, "2025-11-25"): async def close_stream_callback() -> None: @@ -272,22 +548,22 @@ 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, ) - else: - metadata = ServerMessageMetadata(request_context=request) + return ServerMessageMetadata(request_context=request) - return SessionMessage(message, metadata=metadata) + def _transport_context(self, request: Request, *, can_send_request: bool) -> TransportContext: + return TransportContext(kind=STREAMABLE_HTTP_KIND, 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: @@ -300,31 +576,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, @@ -390,19 +641,13 @@ def _create_event_data(self, event_message: EventMessage) -> SSEEvent: return event_data - 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.""" @@ -462,11 +707,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): @@ -532,10 +782,7 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re 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) - + await self._deliver_client_message(request, message) return # Extract protocol version for priming event decision. @@ -547,99 +794,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) - - 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) - 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() + await self._serve_request(scope, request, receive, send, message, protocol_version) - except Exception as err: + except Exception: logger.exception("Error handling POST request") response = self._create_error_response( "Error handling POST request", @@ -647,9 +804,251 @@ 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.""" + assert self._runner is not None, "stateful session was built without a server" + 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: + try: + await runner.on_notify(dctx, message.method, message.params) + except Exception: # a crashing notification handler must not take the session down + logger.exception("notification handler for %r raised", message.method) + 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. + + Stateless mode and JSON-response mode have none: a nested elicitation + or sampling request has no stream to ride, and no POST of the client's + answer could be correlated back to a session. + """ + 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 if one is running, else inline.""" + tg = self._task_group + if tg is not None: + tg.start_soon(fn, *args) + else: + await 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 = str(request_id) + + # 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) + ) + + channel = _MessageChannel(stream_id, self._event_store) + self._channels[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 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=channel, + _request_id=request_id, + message_metadata=self._build_message_metadata(request, request_id, protocol_version), + _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: + 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), + ) + except Exception: + # `serve_inbound` contains handler exceptions itself; anything + # escaping it is a channel write that raised (e.g. a broken + # EventStore), which must cost this request, not the session. + logger.exception("Error handling request %r", request_id) + finally: + # The channel stays registered until the handler is done, so a + # `Last-Event-ID` reconnect can re-attach while it still runs. + channel.finish() + if self._channels.get(stream_id) is channel: + del self._channels[stream_id] + if connection is not None: + await aclose_shielded(connection) + + if self._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. + self._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() + + 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) + 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) + + async def _respond_sse( + self, + scope: Scope, + receive: Receive, + send: Send, + channel: _MessageChannel, + reader: MemoryObjectReceiveStream[EventMessage], + priming_event: SSEEvent | None, + ) -> None: + """Stream the request's channel as this POST's SSE response, until the response frame passes.""" + sse_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) + response = EventSourceResponse( + content=sse_recv, + data_sender_callable=partial( + self._pump_channel, channel, reader, sse_send, priming_event, stop_at_response=True + ), + headers=self._sse_headers(), + ) + try: + await response(scope, receive, send) + finally: + # The client is gone (disconnect or delivered response): detach so + # the handler carries on writing to the store alone. + channel.detach(reader) + await sse_send.aclose() + await sse_recv.aclose() + + async def _pump_channel( + self, + channel: _MessageChannel, + reader: MemoryObjectReceiveStream[EventMessage], + sse_send: MemoryObjectSendStream[SSEEvent], + priming_event: SSEEvent | None, + *, + 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 + 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") + channel.detach(reader) + async def _handle_get_request(self, request: Request, send: Send) -> None: """Handle GET request to establish SSE. @@ -657,10 +1056,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) @@ -676,21 +1071,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, @@ -698,54 +1084,24 @@ 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 + reader = self._standalone.attach() + sse_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=standalone_sse_writer, - headers=headers, + content=sse_recv, + data_sender_callable=partial( + self._pump_channel, self._standalone, reader, sse_send, None, stop_at_response=False + ), + headers=self._sse_headers(), ) - 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) finally: - await sse_stream_writer.aclose() - await sse_stream_reader.aclose() + self._standalone.detach(reader) + await sse_send.aclose() + await sse_recv.aclose() async def _handle_delete_request(self, request: Request, send: Send) -> None: """Handle DELETE requests for explicit session termination.""" @@ -779,27 +1135,17 @@ async def terminate(self) -> None: self._terminated = True 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._channels.values()): + channel.close() + self._channels.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.""" @@ -862,60 +1208,48 @@ 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 + headers = self._sse_headers() # 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) + sse_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) - async def replay_sender(): + async def replay_sender() -> None: try: - async with sse_stream_writer: + async with sse_send: # 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) + await sse_send.send(self._create_event_data(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) + # 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 not stream_id: + return + channel = self._standalone if stream_id == GET_STREAM_KEY else self._channels.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() + 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 + await self._pump_channel(channel, reader, sse_send, priming_event, stop_at_response=True) + finally: + channel.detach(reader) except anyio.ClosedResourceError: # pragma: lax no cover # Expected when close_sse_stream() is called logger.debug("Replay SSE stream closed by close_sse_stream()") @@ -924,7 +1258,7 @@ async def send_event(event_message: EventMessage) -> None: # Create and start EventSourceResponse response = EventSourceResponse( - content=sse_stream_reader, + content=sse_recv, data_sender_callable=replay_sender, headers=headers, ) @@ -934,8 +1268,8 @@ async def send_event(event_message: EventMessage) -> None: except Exception: # pragma: lax no cover logger.exception("Error in replay response") finally: - await sse_stream_writer.aclose() - await sse_stream_reader.aclose() + await sse_send.aclose() + await sse_recv.aclose() except Exception: # pragma: lax no cover logger.exception("Error replaying events") @@ -945,107 +1279,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/tests/server/test_streamable_http_manager.py b/tests/server/test_streamable_http_manager.py index 70440d9d03..fe93655244 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._channels, "Transport should have no active request channels" + 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..b748cec0e7 --- /dev/null +++ b/tests/server/test_streamable_http_transport.py @@ -0,0 +1,74 @@ +"""Failure handling in the StreamableHTTP server transport's per-request dispatch.""" + +import anyio +import pytest +from mcp_types import JSONRPCMessage +from starlette.types import Message, Scope + +from mcp.server import Server +from mcp.server.streamable_http import ( + EventCallback, + EventId, + EventStore, + StreamableHTTPServerTransport, + StreamId, +) + + +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_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._channels == {} # pyright: ignore[reportPrivateUsage] + 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/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index aeef25a278..9369ec21bf 100644 --- a/tests/shared/test_streamable_http.py +++ b/tests/shared/test_streamable_http.py @@ -19,7 +19,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 +39,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 +46,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 +1636,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 +1664,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 +2176,6 @@ 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] + transport.close_standalone_sse_stream() 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 + assert "Error in SSE writer" not in caplog.text From fcf3dae2d790f8100466c20bdef2b3d8e1724daa Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:01:02 +0000 Subject: [PATCH 3/7] Close the coverage gaps in the rewritten transport Removes now-dead code (the attach take-over branch, redundant containment around the notification handler, an unreachable None-connection arm, the channel-closing loop in run() teardown) and adds tests for the behaviours that were untested: terminating a session with a request in flight, client-posted progress for a server-initiated request, containment of a raising event store per request, concurrent POSTs sharing a request id, the channel attach/detach identity guard, and the serve_loop driver. --- src/mcp/server/runner.py | 6 +- src/mcp/server/streamable_http.py | 43 ++- src/mcp/shared/jsonrpc_dispatcher.py | 10 +- tests/server/test_runner.py | 22 ++ .../server/test_streamable_http_transport.py | 283 +++++++++++++++++- 5 files changed, 334 insertions(+), 30 deletions(-) diff --git a/src/mcp/server/runner.py b/src/mcp/server/runner.py index 6f9f7a8f74..75c0633aba 100644 --- a/src/mcp/server/runner.py +++ b/src/mcp/server/runner.py @@ -477,9 +477,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 c37ea2a2a6..3c1978da42 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -224,9 +224,8 @@ def attach(self) -> MemoryObjectReceiveStream[EventMessage]: (the standalone GET stream); re-attaching after a detach is how a `Last-Event-ID` reconnect resumes a live stream. """ + assert self._writer is None, "every attach site checks `attached` first" writer, reader = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) - if self._writer is not None: - self._writer.close() self._writer = writer return reader @@ -262,8 +261,15 @@ def close(self) -> None: self.finished.set() def finish(self) -> None: - """Mark the request as finished even if no terminal frame was recorded.""" + """The request is over: mark it finished and end any response still attached. + + Normally the terminal frame already ended the response; this covers a + request whose terminal write never landed (e.g. the event store raised + mid-stream), so the client sees the stream close rather than hanging. + """ self.finished.set() + if self._writer is not None: + self._detach(self._writer) @dataclass @@ -513,6 +519,9 @@ async def run(self, *, task_status: anyio.abc.TaskStatus[None] = anyio.TASK_STAT 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 @@ -521,15 +530,12 @@ async def run(self, *, task_status: anyio.abc.TaskStatus[None] = anyio.TASK_STAT tg.cancel_scope.cancel() finally: self._task_group = None - # End every still-open response stream and wake anything awaiting a - # client answer; runs on termination and on manager shutdown alike. - for channel in list(self._channels.values()): - channel.close() - self._channels.clear() + # 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() - if self._connection is not None: - await aclose_shielded(self._connection) + await aclose_shielded(connection) def _build_message_metadata( self, request: Request, request_id: RequestId, protocol_version: str @@ -650,7 +656,14 @@ def _sse_headers(self) -> dict[str, str]: } 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 @@ -807,8 +820,8 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re return def _session_runner(self) -> ServerRunner[Any]: - """The stateful session's handler kernel.""" - assert self._runner is not None, "stateful session was built without a server" + """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]: @@ -853,10 +866,10 @@ async def _deliver_client_message( ) 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) - except Exception: # a crashing notification handler must not take the session down - logger.exception("notification handler for %r raised", message.method) finally: if connection is not None: await aclose_shielded(connection) diff --git a/src/mcp/shared/jsonrpc_dispatcher.py b/src/mcp/shared/jsonrpc_dispatcher.py index 8ff1cd0647..b964ad36b2 100644 --- a/src/mcp/shared/jsonrpc_dispatcher.py +++ b/src/mcp/shared/jsonrpc_dispatcher.py @@ -19,7 +19,6 @@ import anyio import anyio.abc from mcp_types import ( - CONNECTION_CLOSED, INTERNAL_ERROR, ErrorData, JSONRPCError, @@ -51,7 +50,7 @@ as_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, @@ -271,10 +270,9 @@ 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._corr.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()") plan = _plan_outbound(_related_request_id, opts) return await self._corr.call( 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_transport.py b/tests/server/test_streamable_http_transport.py index b748cec0e7..a08cb22fc1 100644 --- a/tests/server/test_streamable_http_transport.py +++ b/tests/server/test_streamable_http_transport.py @@ -1,18 +1,44 @@ -"""Failure handling in the StreamableHTTP server transport's per-request dispatch.""" +"""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. +""" import anyio import pytest -from mcp_types import JSONRPCMessage +from httpx2 import EventSource +from mcp_types import ( + CallToolRequestParams, + CallToolResult, + ElicitRequest, + ElicitRequestFormParams, + ElicitResult, + JSONRPCMessage, + JSONRPCRequest, + JSONRPCResponse, + TextContent, +) from starlette.types import Message, Scope -from mcp.server import Server +from mcp.server import Server, ServerRequestContext 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 NoBackChannelError +from mcp.shared.message import ServerMessageMetadata +from mcp.shared.transport_context import TransportContext +from tests.interaction._connect import base_headers, initialize_via_http, mounted_app, parse_sse_messages +from tests.interaction.transports._event_store import SequencedEventStore + +pytestmark = pytest.mark.anyio class _PrimingFailingStore(EventStore): @@ -23,7 +49,21 @@ async def replay_events_after(self, last_event_id: EventId, send_callback: Event raise NotImplementedError -@pytest.mark.anyio +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 == "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. @@ -70,5 +110,236 @@ async def asgi_send(message: Message) -> None: assert transport._channels == {} # pyright: ignore[reportPrivateUsage] 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 + 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_mid_request_costs_only_that_request() -> None: + """A store that raises for a request's stream fails that request cleanly, never the session. + + Neither request 42's result nor the error frame reporting the failure can be stored, so its + stream ends without a terminal frame instead of hanging; request 43 on the same session is + 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) + + async with mounted_app(server, event_store=_StreamFailingStore(), 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 stream carried its priming event and then ended with no terminal frame. + assert parse_sse_messages(failing_events) == [] + (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() + channel.detach() # e.g. close_sse_stream() + fresh_reader = channel.attach() # the client's reconnect re-attached + 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 + + dctx.close() + + await dctx.notify("notifications/message", {"level": "info", "data": "too late"}) + assert not channel.finished.is_set() + with pytest.raises(NoBackChannelError): + await dctx.send_raw_request("ping", None) From d9f77bcf8e89a526fa6c67ac3252bf885e92ba36 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:04:55 +0000 Subject: [PATCH 4/7] Document the streamable HTTP transport change in the migration guide Records the removal of StreamableHTTPServerTransport.connect() (the transport is now driven per request by the session manager) and the behaviours clarified alongside it, and refreshes the docstrings that still described the old serve_loop wiring. --- docs/migration.md | 37 +++++++++++++++++++++++++++++++ src/mcp/server/lowlevel/server.py | 5 +++-- tests/interaction/README.md | 5 ++--- 3 files changed, 42 insertions(+), 5 deletions(-) diff --git a/docs/migration.md b/docs/migration.md index 5eb7659c23..833c06806d 100644 --- a/docs/migration.md +++ b/docs/migration.md @@ -721,6 +721,43 @@ When serving streamable HTTP (stateful or `stateless_http=True`), the server's ` Lifespans that set up process-wide state (connection pools, caches, background tasks) are unaffected — they now run once instead of per session/request. If your lifespan was acquiring per-connection resources, move that acquisition into the handler body; per-connection cleanup belongs on the connection's `exit_stack` (a public way to reach it from high-level `@mcp.tool()` handlers is planned). +### 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): let the manager own transports (mount it or use streamable_http_app()) +session_manager = StreamableHTTPSessionManager(app=server, event_store=..., json_response=...) +``` + +Behaviour clarified in the same change: + +- In JSON-response mode a server-to-client request (sampling, elicitation) now raises + `NoBackChannelError` — a JSON response has no stream to carry the request, and previously the + call would hang waiting for an answer that could never be delivered. +- Stateful streamable-HTTP handlers see `TransportContext(kind="streamable-http")` carrying the + request `headers`; the stream-pair path previously reported `kind="jsonrpc"` with no headers. +- 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. + ### `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/tests/interaction/README.md b/tests/interaction/README.md index 3060a240c1..49188b1366 100644 --- a/tests/interaction/README.md +++ b/tests/interaction/README.md @@ -279,9 +279,8 @@ 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 SSE-writer and replay +`except` handlers 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. From edb0abd7b0a43a9dc7115de1b6c2e0f3bae3ed17 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:15:51 +0000 Subject: [PATCH 5/7] Inline the streamable-http transport-context kind literal --- src/mcp/server/streamable_http.py | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 3c1978da42..8cfadf6312 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -90,9 +90,6 @@ # Pattern ensures entire string contains only valid characters by using ^ and $ anchors SESSION_ID_PATTERN = re.compile(r"^[\x21-\x7E]+$") -# Streamable HTTP transport kind for `TransportContext.kind`. -STREAMABLE_HTTP_KIND = "streamable-http" - # Type aliases StreamId = str EventId = str @@ -562,7 +559,7 @@ async def close_standalone_stream_callback() -> None: return ServerMessageMetadata(request_context=request) def _transport_context(self, request: Request, *, can_send_request: bool) -> TransportContext: - return TransportContext(kind=STREAMABLE_HTTP_KIND, can_send_request=can_send_request, headers=request.headers) + 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. From 1e3be94d4a4dc68e413d86e547f36aeea3cbc328 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 20:48:49 +0000 Subject: [PATCH 6/7] Fix defects found in review of the streamable HTTP rewrite - Serialize store-then-forward per channel: concurrent writers on one channel (the standalone GET stream) could put frames on the wire out of event-store order, breaking Last-Event-ID resumption. A per-channel lock restores the store-order == wire-order invariant the serial router used to provide. - Refuse a request that arrives across session termination instead of running it: dispatch now branches on the transport's identity (stateful vs stateless) rather than on session-task presence, re-checks liveness after the request's awaits, and answers 404 for an ended session; queued work for an ended session is dropped. - Contain event-store failures at the channel: a raising store_event was reaching the correlator as a handler error and leaking its text onto the wire; write() now never raises (a broken store ends that stream), which also protects the courtesy-cancel write. - Close the replay reader when priming fails; drop dead exception arms in the SSE pump. - Correct the migration note's JSON-mode claim to request-scoped requests, drop a claim handlers cannot observe, fix a vacuous test assertion, and add regression tests for each fix. --- docs/migration.md | 10 +- src/mcp/server/streamable_http.py | 148 +++++++++------ tests/interaction/README.md | 5 +- .../server/test_streamable_http_transport.py | 172 +++++++++++++++++- tests/shared/test_streamable_http.py | 4 +- 5 files changed, 271 insertions(+), 68 deletions(-) diff --git a/docs/migration.md b/docs/migration.md index 833c06806d..e0fe58f66c 100644 --- a/docs/migration.md +++ b/docs/migration.md @@ -750,11 +750,11 @@ session_manager = StreamableHTTPSessionManager(app=server, event_store=..., json Behaviour clarified in the same change: -- In JSON-response mode a server-to-client request (sampling, elicitation) now raises - `NoBackChannelError` — a JSON response has no stream to carry the request, and previously the - call would hang waiting for an answer that could never be delivered. -- Stateful streamable-HTTP handlers see `TransportContext(kind="streamable-http")` carrying the - request `headers`; the stream-pair path previously reported `kind="jsonrpc"` with no headers. +- 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. diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 8cfadf6312..4feb2af66a 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -175,6 +175,10 @@ def __init__(self, stream_id: StreamId, event_store: EventStore | None) -> None: self._event_store = event_store self._writer: MemoryObjectSendStream[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() @@ -186,33 +190,41 @@ def attached(self) -> bool: return self._writer is not None async def write(self, message: JSONRPCMessage) -> None: - """Store-then-forward one outbound message. Never raises for a dropped connection.""" - if self._closed: - logger.debug("dropped message on closed stream %s", self.stream_id) - return - # 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: - event_id = await self._event_store.store_event(self.stream_id, message) - 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.""" - ) - return - try: - await writer.send(EventMessage(message, event_id)) - except (anyio.BrokenResourceError, anyio.ClosedResourceError): - # The SSE response went away between the attach check and the send. - self._detach(writer) + """Store-then-forward one outbound message. Never raises.""" + async with self._write_lock: + if self._closed: + logger.debug("dropped message on closed stream %s", self.stream_id) + return + # 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 must not leak its exception into the caller's + # write (nor its text onto the wire): this stream is over. + logger.exception("EventStore.store_event failed for stream %s; closing the stream", self.stream_id) + self.close() + return + 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.""" + ) + return + try: + await writer.send(EventMessage(message, event_id)) + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + # The SSE response went away between the attach check and the send. + self._detach(writer) def attach(self) -> MemoryObjectReceiveStream[EventMessage]: """Attach a fresh SSE response and return the reader it drains. @@ -877,19 +889,27 @@ async def _run_notification() -> None: def _can_send_request(self) -> bool: """Whether a request handler on this transport has a back-channel for server-to-client requests. - Stateless mode and JSON-response mode have none: a nested elicitation - or sampling request has no stream to ride, and no POST of the client's - answer could be correlated back to a session. + 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 if one is running, else inline.""" - tg = self._task_group - if tg is not None: - tg.start_soon(fn, *args) - else: + """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 self._terminated: + logger.debug("dropped work for ended session %s", self.mcp_session_id) + return + tg.start_soon(fn, *args) async def _serve_request( self, @@ -912,18 +932,32 @@ async def _serve_request( 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. + stateful = self.mcp_session_id is not None + session_task_group = self._task_group + if stateful and (self._terminated 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 + channel = _MessageChannel(stream_id, self._event_store) self._channels[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 self.mcp_session_id is None: - runner = self._stateless_runner(request) - connection = runner.connection - else: + 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, @@ -937,6 +971,8 @@ async def _serve_request( 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, @@ -945,11 +981,6 @@ async def _run_handler() -> None: write_result=partial(self._write_result, channel, request_id), write_error=partial(self._write_error, channel, request_id), ) - except Exception: - # `serve_inbound` contains handler exceptions itself; anything - # escaping it is a channel write that raised (e.g. a broken - # EventStore), which must cost this request, not the session. - logger.exception("Error handling request %r", request_id) finally: # The channel stays registered until the handler is done, so a # `Last-Event-ID` reconnect can re-attach while it still runs. @@ -959,11 +990,11 @@ async def _run_handler() -> None: if connection is not None: await aclose_shielded(connection) - if self._task_group is not None: + 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. - self._task_group.start_soon(_run_handler) + session_task_group.start_soon(_run_handler) if reader is None: await self._respond_json(scope, receive, send, channel) else: @@ -1051,10 +1082,6 @@ async def _pump_channel( await sse_send.send(self._create_event_data(event_message)) if stop_at_response and 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") channel.detach(reader) @@ -1256,14 +1283,23 @@ async def send_event(event_message: EventMessage) -> None: # stream is re-registered. priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) - # Forward messages to SSE - await self._pump_channel(channel, reader, sse_send, priming_event, stop_at_response=True) + # 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, + sse_send, + priming_event, + stop_at_response=stream_id != GET_STREAM_KEY, + ) finally: channel.detach(reader) - 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 + # 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 diff --git a/tests/interaction/README.md b/tests/interaction/README.md index 49188b1366..282c0e0d19 100644 --- a/tests/interaction/README.md +++ b/tests/interaction/README.md @@ -279,8 +279,9 @@ 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 SSE-writer and replay -`except` handlers in `server/streamable_http.py`. `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 standalone-GET and replay `response(...)` calls 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_streamable_http_transport.py b/tests/server/test_streamable_http_transport.py index a08cb22fc1..df5d8328ad 100644 --- a/tests/server/test_streamable_http_transport.py +++ b/tests/server/test_streamable_http_transport.py @@ -5,20 +5,26 @@ that the transport-agnostic interaction matrix does not reach. """ +from unittest.mock import MagicMock + import anyio +import anyio.lowlevel import pytest from httpx2 import EventSource from mcp_types import ( + INVALID_REQUEST, CallToolRequestParams, CallToolResult, ElicitRequest, ElicitRequestFormParams, ElicitResult, JSONRPCMessage, + JSONRPCNotification, JSONRPCRequest, JSONRPCResponse, TextContent, ) +from starlette.requests import Request from starlette.types import Message, Scope from mcp.server import Server, ServerRequestContext @@ -32,10 +38,16 @@ _MessageChannel, # pyright: ignore[reportPrivateUsage] ) from mcp.shared._correlation import RequestCorrelator -from mcp.shared.exceptions import NoBackChannelError +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, initialize_via_http, mounted_app, parse_sse_messages +from tests.interaction._connect import ( + base_headers, + connect_over_streamable_http, + initialize_via_http, + mounted_app, + parse_sse_messages, +) from tests.interaction.transports._event_store import SequencedEventStore pytestmark = pytest.mark.anyio @@ -336,10 +348,164 @@ async def test_a_closed_request_context_drops_notifications_and_refuses_requests _request_id=1, ) assert dctx.can_send_request + reader = channel.attach() dctx.close() await dctx.notify("notifications/message", {"level": "info", "data": "too late"}) - assert not channel.finished.is_set() + 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() + 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 == "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") + ) diff --git a/tests/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index 9369ec21bf..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 @@ -2177,5 +2178,4 @@ async def message_handler(message: IncomingMessage) -> None: # Tear the standalone stream down while the writer is parked on it. (transport,) = session_manager._server_instances.values() # pyright: ignore[reportPrivateUsage] transport.close_standalone_sse_stream() - assert "Error in standalone SSE writer" not in caplog.text - assert "Error in SSE writer" not in caplog.text + assert [r for r in caplog.records if r.name == "mcp.server.streamable_http" and r.levelno >= logging.ERROR] == [] From 47e156d5609ac8bce1d0993c1c9fe10827bbb257 Mon Sep 17 00:00:00 2001 From: Max Isbey <224885523+maxisbey@users.noreply.github.com> Date: Mon, 27 Jul 2026 23:19:39 +0000 Subject: [PATCH 7/7] Rework channel liveness, stream ids, and init ordering after review The second review round found several defects that traced back to a few structural gaps rather than isolated bugs; this addresses the gaps. - Channel delivery is now a value: write() returns whether the message reached the event store or an attached response, an EventStore failure is contained inside the write (resumability degrades, the stream survives, no store text on the wire), and attach() returns None on a dead channel so nothing can stream from one. A server-to-client request that cannot reach any client fails the caller with CONNECTION_CLOSED instead of parking it, and the JSON-response body derives from the channel's recorded outcome (result, terminated-404, or 500) rather than a "cannot happen" branch. - Stream ids handed to the EventStore are minted by the transport in a session-scoped namespace, so a client-chosen request id can no longer name the standalone GET stream and two sessions on one store can no longer replay each other's frames. - One SSE-response runner owns the pump, error containment, and cleanup for the POST, GET, and replay responses (the POST path had lost the guard its siblings kept). - A session-level gate holds requests that arrive while an initialize is still being served until the handshake commits, restoring the ordering the stream-pair driver's parked read loop used to guarantee. - The POSTed client message is delivered even when the 202 could not be written back, and the correlator marks the single site where the cancelled-request answer policy lives for every transport. Adds regression tests for each of the above. --- docs/migration.md | 20 +- src/mcp/server/runner.py | 7 +- src/mcp/server/streamable_http.py | 370 ++++++++++------ src/mcp/shared/_correlation.py | 5 + tests/interaction/README.md | 3 +- tests/server/test_streamable_http_manager.py | 2 +- .../server/test_streamable_http_transport.py | 395 +++++++++++++++++- 7 files changed, 659 insertions(+), 143 deletions(-) diff --git a/docs/migration.md b/docs/migration.md index e0fe58f66c..cbccdf5501 100644 --- a/docs/migration.md +++ b/docs/migration.md @@ -744,10 +744,14 @@ 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): let the manager own transports (mount it or use streamable_http_app()) -session_manager = StreamableHTTPSessionManager(app=server, event_store=..., json_response=...) +# 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 @@ -757,6 +761,18 @@ Behaviour clarified in the same change: `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 diff --git a/src/mcp/server/runner.py b/src/mcp/server/runner.py index 75c0633aba..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 diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 4feb2af66a..7a2fe9540e 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -22,11 +22,12 @@ import logging import re from abc import ABC, abstractmethod -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import Awaitable, Callable, Coroutine, Mapping from dataclasses import dataclass, field from functools import partial from http import HTTPStatus from typing import TYPE_CHECKING, Any, Final +from uuid import uuid4 import anyio import anyio.abc @@ -155,7 +156,9 @@ 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 @@ -165,15 +168,17 @@ class _MessageChannel: 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. With no - response attached the message is only stored (or, without a store, - dropped with a debug log) - the client can reconnect and replay. + 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 @@ -182,19 +187,24 @@ def __init__(self, stream_id: StreamId, event_store: EventStore | None) -> None: 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 closed by termination).""" + """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) -> None: - """Store-then-forward one outbound message. Never raises.""" + 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 + 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 @@ -203,12 +213,12 @@ async def write(self, message: JSONRPCMessage) -> None: try: event_id = await self._event_store.store_event(self.stream_id, message) except Exception: - # A broken store must not leak its exception into the caller's - # write (nor its text onto the wire): this stream is over. - logger.exception("EventStore.store_event failed for stream %s; closing the stream", self.stream_id) - self.close() - return - logger.debug(f"Stored {event_id} from {self.stream_id}") + # 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() @@ -219,66 +229,56 @@ async def write(self, message: JSONRPCMessage) -> None: for message. Still processing message as the client might reconnect and replay.""" ) - return + # 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 SSE response went away between the attach check and the send. - self._detach(writer) + # 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]: + def attach(self) -> MemoryObjectReceiveStream[EventMessage] | None: """Attach a fresh SSE response and return the reader it drains. - 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. + 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 = writer + 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 one attachment: a stale response ending - must not knock a newer (resumed) attachment off the channel. + `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. """ - writer = self._writer - if writer is None: + if self._writer is None: return - if reader is not None and not self._is_writer_of(writer, reader): + if reader is not None and reader is not self._reader: return - self._detach(writer) - - def _is_writer_of( - self, writer: MemoryObjectSendStream[EventMessage], reader: MemoryObjectReceiveStream[EventMessage] - ) -> bool: - # A memory-object stream pair shares one state object. - return getattr(writer, "_state", None) is getattr(reader, "_state", None) - - def _detach(self, writer: MemoryObjectSendStream[EventMessage]) -> None: - if self._writer is writer: - self._writer = None - writer.close() + self._writer.close() + self._writer, self._reader = None, None def close(self) -> None: - """End the channel outright (session termination): detach and refuse further writes.""" - self._closed = True - if self._writer is not None: - self._detach(self._writer) - self.finished.set() - - def finish(self) -> None: - """The request is over: mark it finished and end any response still attached. + """The stream is over: detach any response and drop every further write. - Normally the terminal frame already ended the response; this covers a - request whose terminal write never landed (e.g. the event store raised - mid-stream), so the client sees the stream close rather than hanging. + 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() - if self._writer is not None: - self._detach(self._writer) + self.detach() @dataclass @@ -381,6 +381,13 @@ async def _call_over_channel( """ 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})) @@ -388,7 +395,7 @@ async def send_cancel(request_id: RequestId, reason: str) -> None: method, params, opts, - write_request=channel.write, + write_request=write_request, send_cancel=send_cancel, cancel_on_abandon=opts.get("cancel_on_abandon", True), ) @@ -455,15 +462,27 @@ def __init__( # 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(GET_STREAM_KEY, event_store) - # In-flight request streams, keyed by stream id (str of the request id), - # so `close_sse_stream()` and `Last-Event-ID` replay can find them. - self._channels: dict[StreamId, _MessageChannel] = {} + 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. @@ -479,6 +498,19 @@ 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. @@ -497,8 +529,7 @@ def close_sse_stream(self, request_id: RequestId) -> None: Requires event_store to be configured for events to be stored during the disconnect. """ - stream_id = str(request_id) - channel = self._standalone if stream_id == GET_STREAM_KEY else self._channels.get(stream_id) + channel = self._streams.get(self._request_stream_id(request_id)) if channel is not None: channel.detach() @@ -535,7 +566,12 @@ async def run(self, *, task_status: anyio.abc.TaskStatus[None] = anyio.TASK_STAT async with anyio.create_task_group() as tg: self._task_group = tg task_status.started() - await self._closed_event.wait() + 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 @@ -801,10 +837,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 - await self._deliver_client_message(request, 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. @@ -906,7 +945,7 @@ async def _spawn_or_run(self, fn: Callable[..., Awaitable[None]], *args: Any) -> await fn(*args) return tg = self._task_group - if tg is None or self._terminated: + 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) @@ -922,7 +961,52 @@ async def _serve_request( ) -> None: """Dispatch one POSTed JSON-RPC request and stream its response.""" request_id = message.id - stream_id = str(request_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 @@ -936,18 +1020,17 @@ async def _serve_request( # 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. - stateful = self.mcp_session_id is not None session_task_group = self._task_group - if stateful and (self._terminated or session_task_group is None): + 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 + return False channel = _MessageChannel(stream_id, self._event_store) - self._channels[stream_id] = channel + 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() @@ -982,11 +1065,13 @@ async def _run_handler() -> None: write_error=partial(self._write_error, channel, request_id), ) 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.finish() - if self._channels.get(stream_id) is channel: - del self._channels[stream_id] + channel.close() + if self._streams.get(stream_id) is channel: + del self._streams[stream_id] if connection is not None: await aclose_shielded(connection) @@ -1010,6 +1095,7 @@ async def _run_handler() -> None: else: await self._respond_sse(scope, receive, send, channel, reader, priming_event) tg.cancel_scope.cancel() + return True 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)) @@ -1023,8 +1109,16 @@ async def _respond_json(self, scope: Scope, receive: Receive, send: Send, channe response_message = channel.terminal if response_message is not None: response = self._create_json_response(response_message) - else: # pragma: no cover - # This shouldn't happen in normal operation + 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", @@ -1038,24 +1132,43 @@ async def _respond_sse( receive: Receive, send: Send, channel: _MessageChannel, - reader: MemoryObjectReceiveStream[EventMessage], + 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( - self._pump_channel, channel, reader, sse_send, priming_event, stop_at_response=True - ), + 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: - # The client is gone (disconnect or delivered response): detach so - # the handler carries on writing to the store alone. - channel.detach(reader) await sse_send.aclose() await sse_recv.aclose() @@ -1063,8 +1176,8 @@ async def _pump_channel( self, channel: _MessageChannel, reader: MemoryObjectReceiveStream[EventMessage], - sse_send: MemoryObjectSendStream[SSEEvent], priming_event: SSEEvent | None, + sse_send: MemoryObjectSendStream[SSEEvent], *, stop_at_response: bool, ) -> None: @@ -1122,23 +1235,24 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: return reader = self._standalone.attach() - sse_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) - response = EventSourceResponse( - content=sse_recv, - data_sender_callable=partial( - self._pump_channel, self._standalone, reader, sse_send, None, stop_at_response=False - ), - headers=self._sse_headers(), - ) + 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._run_sse_response( + request.scope, + request.receive, + send, + partial(self._pump_channel, self._standalone, reader, None, stop_at_response=False), + ) finally: self._standalone.detach(reader) - await sse_send.aclose() - await sse_recv.aclose() async def _handle_delete_request(self, request: Request, send: Send) -> None: """Handle DELETE requests for explicit session termination.""" @@ -1170,13 +1284,14 @@ async def terminate(self) -> None: """ self._terminated = True + self._accepting = False logger.info(f"Terminating session: {self.mcp_session_id}") # Close every open response stream, wake anything awaiting a # client answer, and cancel in-flight handlers. - for channel in list(self._channels.values()): + for channel in list(self._streams.values()): channel.close() - self._channels.clear() + self._streams.clear() self._standalone.close() self._corr.close() self._corr.cancel_all_inbound() @@ -1245,30 +1360,39 @@ async def _replay_events(self, last_event_id: str, request: Request, send: Send) return # pragma: no cover try: - headers = self._sse_headers() - # 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_send, sse_recv = anyio.create_memory_object_stream[SSEEvent](0) - - async def replay_sender() -> None: + async def replay_then_tail(sse_send: MemoryObjectSendStream[SSEEvent]) -> None: try: async with sse_send: - # Define an async callback for sending events - async def send_event(event_message: EventMessage) -> None: - await sse_send.send(self._create_event_data(event_message)) + # 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) + 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 not stream_id: - return - channel = self._standalone if stream_id == GET_STREAM_KEY else self._channels.get(stream_id) + 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 @@ -1278,6 +1402,10 @@ async def send_event(event_message: EventMessage) -> None: # 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. @@ -1289,9 +1417,9 @@ async def send_event(event_message: EventMessage) -> None: await self._pump_channel( channel, reader, - sse_send, priming_event, - stop_at_response=stream_id != GET_STREAM_KEY, + sse_send, + stop_at_response=channel is not self._standalone, ) finally: channel.detach(reader) @@ -1303,19 +1431,7 @@ async def send_event(event_message: EventMessage) -> None: logger.exception("Error in replay sender") # Create and start EventSourceResponse - response = EventSourceResponse( - content=sse_recv, - 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_send.aclose() - await sse_recv.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") diff --git a/src/mcp/shared/_correlation.py b/src/mcp/shared/_correlation.py index 008d740584..8e2793e8e2 100644 --- a/src/mcp/shared/_correlation.py +++ b/src/mcp/shared/_correlation.py @@ -448,6 +448,11 @@ async def serve_inbound( # result write above did not happen - no double response. # TODO(L38): spec says SHOULD NOT respond after cancel; # the existing server always has, so match that for now. + # This is the single site every transport shares for the + # cancelled-request answer policy; change it here for all + # of them (a transport whose response stream must still end + # sees this write, so a suppressed answer needs the stream + # terminated another way). answer_write_started = True await write_error(ErrorData(code=0, message="Request cancelled")) except anyio.get_cancelled_exc_class(): diff --git a/tests/interaction/README.md b/tests/interaction/README.md index 282c0e0d19..17a4a7c332 100644 --- a/tests/interaction/README.md +++ b/tests/interaction/README.md @@ -280,7 +280,8 @@ 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 `except Exception` arms -around the standalone-GET and replay `response(...)` calls in `server/streamable_http.py`. +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 diff --git a/tests/server/test_streamable_http_manager.py b/tests/server/test_streamable_http_manager.py index fe93655244..17d216c887 100644 --- a/tests/server/test_streamable_http_manager.py +++ b/tests/server/test_streamable_http_manager.py @@ -423,7 +423,7 @@ async def mock_receive(): assert transport._terminated, "Transport should be terminated after stateless request" # Verify internal state is cleaned up: no request streams left open. - assert not transport._channels, "Transport should have no active request channels" + assert not transport._streams, "Transport should have no active request streams" assert not transport._standalone.attached, "Transport should have no standalone stream attached" diff --git a/tests/server/test_streamable_http_transport.py b/tests/server/test_streamable_http_transport.py index df5d8328ad..a129988314 100644 --- a/tests/server/test_streamable_http_transport.py +++ b/tests/server/test_streamable_http_transport.py @@ -5,6 +5,7 @@ that the transport-agnostic interaction matrix does not reach. """ +from typing import Any from unittest.mock import MagicMock import anyio @@ -12,12 +13,14 @@ import pytest from httpx2 import EventSource from mcp_types import ( + CONNECTION_CLOSED, INVALID_REQUEST, CallToolRequestParams, CallToolResult, ElicitRequest, ElicitRequestFormParams, ElicitResult, + JSONRPCError, JSONRPCMessage, JSONRPCNotification, JSONRPCRequest, @@ -28,6 +31,7 @@ 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, @@ -44,6 +48,7 @@ from tests.interaction._connect import ( base_headers, connect_over_streamable_http, + initialize_body, initialize_via_http, mounted_app, parse_sse_messages, @@ -65,7 +70,7 @@ 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 == "42" and message is not None: + if stream_id.endswith(":request:42") and message is not None: raise RuntimeError("backend fell over") return await super().store_event(stream_id, message) @@ -119,7 +124,7 @@ async def asgi_send(message: Message) -> None: with anyio.fail_after(5): await transport.handle_request(scope, receive, asgi_send) - assert transport._channels == {} # pyright: ignore[reportPrivateUsage] + 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") @@ -224,20 +229,21 @@ async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> assert reports == [(0.5, 1.0, "half")] -async def test_an_event_store_failure_mid_request_costs_only_that_request() -> None: - """A store that raises for a request's stream fails that request cleanly, never the session. +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. - Neither request 42's result nor the error frame reporting the failure can be stored, so its - stream ends without a terminal frame instead of hanging; request 43 on the same session is - served normally. + 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=_StreamFailingStore(), retry_interval=0) as (http, _): + 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( @@ -251,8 +257,14 @@ async def call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> assert healthy.status_code == 200 healthy_events = [event async for event in EventSource(healthy)] - # Request 42's stream carried its priming event and then ended with no terminal frame. - assert parse_sse_messages(failing_events) == [] + # 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"}] @@ -305,8 +317,10 @@ def test_detaching_a_stale_attachment_does_not_evict_the_newer_one() -> None: 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) @@ -349,6 +363,7 @@ async def test_a_closed_request_context_drops_notifications_and_refuses_requests ) assert dctx.can_send_request reader = channel.attach() + assert reader is not None dctx.close() @@ -386,6 +401,7 @@ async def test_concurrent_writes_reach_the_wire_in_event_store_order() -> None: """ 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") @@ -410,7 +426,7 @@ def __init__(self) -> None: self.release = anyio.Event() async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: - if stream_id == "42" and message is None: + 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) @@ -509,3 +525,360 @@ async def test_a_session_bound_transport_drops_notifications_once_its_session_ha 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