Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 36 additions & 11 deletions src/mcp/client/streamable_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -422,6 +422,7 @@ async def _handle_sse_response(
assert isinstance(ctx.session_message.message, JSONRPCRequest)
original_request_id = ctx.session_message.message.id

stream_fault: Exception | None = None
try:
event_source = EventSource(response)
async for sse in event_source: # pragma: no branch
Expand All @@ -444,19 +445,27 @@ async def _handle_sse_response(
if is_complete:
await response.aclose()
return # Normal completion, no reconnect needed
except Exception:
except Exception as exc:
stream_fault = exc
logger.debug("SSE stream ended", exc_info=True) # pragma: lax no cover

if stream_fault is not None:
# Surface the transport fault to the session's message_handler as an
# Exception item (see IncomingMessage); the waiter itself is resolved
# below with a synthesized error carrying the same exception.
await self._forward_stream_fault(ctx.read_stream_writer, stream_fault)

# Stream ended without response - reconnect if we received an event with ID
if last_event_id is not None:
logger.info("SSE stream disconnected, reconnecting...")
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms)
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms, last_exc=stream_fault)
else:
# Not resumable: resolve the waiter, else a listen stream's consumer
# would hang forever instead of learning the subscription is lost.
await self._resolve_abandoned_request(
ctx.read_stream_writer, original_request_id, "SSE stream ended without a response"
)
message = "SSE stream ended without a response"
if stream_fault is not None:
message = f"{message}: {stream_fault!r}"
await self._resolve_abandoned_request(ctx.read_stream_writer, original_request_id, message)

async def _resolve_abandoned_request(
self, read_stream_writer: StreamWriter, request_id: RequestId, message: str, *, code: int = CONNECTION_CLOSED
Expand All @@ -472,12 +481,27 @@ async def _resolve_abandoned_request(
except (anyio.BrokenResourceError, anyio.ClosedResourceError):
logger.debug("read stream closed before request %r could be resolved", request_id)

async def _forward_stream_fault(self, read_stream_writer: StreamWriter, exc: Exception) -> None:
"""Forward a transport-level fault to the session as an Exception item.

The dispatcher routes Exception items to the session's message_handler via
its on_stream_exception observer (see IncomingMessage); pending request
waiters are resolved separately by the caller with a synthesized error
carrying the same exception.
"""
try:
await read_stream_writer.send(exc)
except (anyio.BrokenResourceError, anyio.ClosedResourceError):
logger.debug("read stream closed before transport fault %r could be forwarded", exc)

async def _handle_reconnection(
self,
ctx: RequestContext,
last_event_id: str,
retry_interval_ms: int | None = None,
attempt: int = 0,
*,
last_exc: Exception | None = None,
) -> None:
"""Reconnect with Last-Event-ID to resume stream after server disconnect."""
# Only requests reconnect: every caller arrives from a request's response stream.
Expand All @@ -488,9 +512,10 @@ async def _handle_reconnection(
# Resolve on give-up: a request with no read timeout (a listen
# stream) would otherwise hang its caller forever.
logger.debug(f"Max reconnection attempts ({MAX_RECONNECTION_ATTEMPTS}) exceeded")
await self._resolve_abandoned_request(
ctx.read_stream_writer, original_request_id, "SSE stream ended and reconnection attempts were exhausted"
)
message = "SSE stream ended and reconnection attempts were exhausted"
if last_exc is not None:
message = f"{message}: {last_exc!r}"
await self._resolve_abandoned_request(ctx.read_stream_writer, original_request_id, message)
return

# Always wait - use server value or default
Expand Down Expand Up @@ -527,11 +552,11 @@ async def _handle_reconnection(

# Stream ended again without response - reconnect again (reset attempt counter)
logger.info("SSE stream disconnected, reconnecting...")
await self._handle_reconnection(ctx, reconnect_last_event_id, reconnect_retry_ms, 0)
except Exception as e: # pragma: no cover
await self._handle_reconnection(ctx, reconnect_last_event_id, reconnect_retry_ms, 0, last_exc=last_exc)
except Exception as e:
logger.debug(f"Reconnection failed: {e}")
# Try to reconnect again if we still have an event ID
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms, attempt + 1)
await self._handle_reconnection(ctx, last_event_id, retry_interval_ms, attempt + 1, last_exc=e)

async def post_writer(
self,
Expand Down
96 changes: 94 additions & 2 deletions tests/client/test_streamable_http.py
Original file line number Diff line number Diff line change
Expand Up @@ -611,7 +611,8 @@ async def aclose(self) -> None:
@pytest.mark.anyio
async def test_a_non_resumable_sse_drop_resolves_the_request_with_an_error() -> None:
"""A per-request SSE stream that dies having carried no event ids can never deliver its
response; the transport resolves the waiter with CONNECTION_CLOSED instead of hanging forever."""
response; the transport surfaces the fault as an Exception item and resolves the waiter
with CONNECTION_CLOSED carrying the exception instead of hanging forever."""
dying = _DyingSSEStream()

def handler(request: httpx2.Request) -> httpx2.Response:
Expand All @@ -625,11 +626,84 @@ def handler(request: httpx2.Request) -> httpx2.Response:
await write.send(
SessionMessage(JSONRPCRequest(jsonrpc="2.0", id="listen-1", method="subscriptions/listen", params={}))
)
fault = await read.receive()
reply = await read.receive()
assert isinstance(fault, httpx2.ReadError)
assert isinstance(reply, SessionMessage)
assert isinstance(reply.message, JSONRPCError)
assert reply.message.id == "listen-1"
assert reply.message.error.code == CONNECTION_CLOSED
assert "connection reset" in reply.message.error.message


class _DyingIdSSEStream(httpx2.AsyncByteStream):
"""Yields an event id with a zero retry, then dies every time it is iterated."""

async def __aiter__(self) -> AsyncIterator[bytes]:
yield b"id: evt-7\nretry: 0\n\n"
raise httpx2.ReadError("connection reset")

async def aclose(self) -> None:
pass


@pytest.mark.anyio
async def test_an_id_bearing_stream_that_dies_surfaces_the_fault_and_resolves_the_request() -> None:
"""A resumable stream whose reconnection attempts all fail forwards the transport fault as
an Exception item and resolves the waiter with CONNECTION_CLOSED carrying the last failure
instead of hanging forever."""

def handler(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(200, headers={"content-type": "text/event-stream"}, stream=_DyingIdSSEStream())

with anyio.fail_after(5):
async with (
httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http,
streamable_http_client("http://test/mcp", http_client=http) as (read, write),
):
await write.send(
SessionMessage(JSONRPCRequest(jsonrpc="2.0", id="listen-1", method="subscriptions/listen", params={}))
)
fault = await read.receive()
reply = await read.receive()
assert isinstance(fault, httpx2.ReadError)
assert isinstance(reply, SessionMessage)
assert isinstance(reply.message, JSONRPCError)
assert reply.message.id == "listen-1"
assert reply.message.error.code == CONNECTION_CLOSED
assert "connection reset" in reply.message.error.message


class _EmptySSEStream(httpx2.AsyncByteStream):
"""An SSE stream that ends cleanly before yielding any events."""

async def __aiter__(self) -> AsyncIterator[bytes]:
return
yield # pragma: no cover


@pytest.mark.anyio
async def test_a_stream_that_ends_cleanly_without_a_response_keeps_the_plain_message() -> None:
"""A per-request SSE stream that ends cleanly (no events, no fault) resolves the
waiter with the plain CONNECTION_CLOSED message."""

def handler(request: httpx2.Request) -> httpx2.Response:
return httpx2.Response(200, headers={"content-type": "text/event-stream"}, stream=_EmptySSEStream())

with anyio.fail_after(5):
async with (
httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http,
streamable_http_client("http://test/mcp", http_client=http) as (read, write),
):
await write.send(
SessionMessage(JSONRPCRequest(jsonrpc="2.0", id="listen-1", method="subscriptions/listen", params={}))
)
reply = await read.receive()
assert isinstance(reply, SessionMessage)
assert isinstance(reply.message, JSONRPCError)
assert reply.message.id == "listen-1"
assert reply.message.error.code == CONNECTION_CLOSED
assert reply.message.error.message == "SSE stream ended without a response"


class _DeliverOnCommandSSEStream(httpx2.AsyncByteStream):
Expand Down Expand Up @@ -725,13 +799,18 @@ async def test_exhausted_reconnection_attempts_resolve_the_request_with_an_error
async with httpx2.AsyncClient() as http:
with anyio.fail_after(5):
await transport._handle_reconnection( # pyright: ignore[reportPrivateUsage]
_abandoned_request_context(http, send), "evt-7", None, MAX_RECONNECTION_ATTEMPTS
_abandoned_request_context(http, send),
"evt-7",
None,
MAX_RECONNECTION_ATTEMPTS,
last_exc=httpx2.ReadError("connection reset"),
)
reply = await receive.receive()
assert isinstance(reply, SessionMessage)
assert isinstance(reply.message, JSONRPCError)
assert reply.message.id == "listen-1"
assert reply.message.error.code == CONNECTION_CLOSED
assert "connection reset" in reply.message.error.message
send.close()
receive.close()

Expand All @@ -748,3 +827,16 @@ async def test_resolving_an_abandoned_request_after_the_reader_closed_is_contain
_abandoned_request_context(http, send), "evt-7", None, MAX_RECONNECTION_ATTEMPTS
)
send.close()


@pytest.mark.anyio
async def test_forwarding_a_stream_fault_after_the_reader_closed_is_contained() -> None:
"""Teardown race: forwarding a fault after the reader closed is best-effort and must not crash."""
transport = StreamableHTTPTransport("http://test/mcp")
send, receive = create_context_streams[SessionMessage | Exception](1)
receive.close()
with anyio.fail_after(5):
await transport._forward_stream_fault( # pyright: ignore[reportPrivateUsage]
send, httpx2.ReadError("connection reset")
)
send.close()
Loading