Skip to content

Commit f7ef5f8

Browse files
committed
Close client SSE iterators explicitly when supported
1 parent e1538c7 commit f7ef5f8

7 files changed

Lines changed: 195 additions & 73 deletions

File tree

‎docs/client/transports.md‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,11 @@ That is the whole production client. `Client` wraps the URL in `streamable_http_
2626

2727
Nothing was resolved, fetched or spawned when you wrote `Client("http://...")`. That line is free.
2828

29+
!!! note "Event-stream cleanup"
30+
Streamable HTTP and legacy SSE close each event iterator when they stop reading, if it
31+
exposes `aclose()`. This runs its cleanup on early return or cancellation instead of leaving
32+
it to garbage collection. HTTPX2 remains responsible for nested response iterators.
33+
2934
### Bring your own `httpx2.AsyncClient`
3035

3136
The moment you need an `Authorization` header, a cookie, a proxy, mTLS, or a different timeout, build the `httpx2.AsyncClient` yourself and hand it to `streamable_http_client`:

‎src/mcp/client/sse.py‎

Lines changed: 44 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
McpHttpClientFactory,
1717
create_mcp_http_client,
1818
request_within_origin,
19+
sse_events,
1920
sse_within_origin,
2021
)
2122
from mcp.shared.message import SessionMessage
@@ -76,48 +77,49 @@ async def sse_client(
7677

7778
async def sse_reader(task_status: TaskStatus[str] = anyio.TASK_STATUS_IGNORED):
7879
try:
79-
async for sse in event_source: # pragma: no branch
80-
logger.debug(f"Received SSE event: {sse.event}")
81-
match sse.event:
82-
case "endpoint":
83-
endpoint_url = urljoin(url, sse.data)
84-
logger.debug(f"Received endpoint URL: {endpoint_url}")
85-
86-
url_parsed = urlparse(url)
87-
endpoint_parsed = urlparse(endpoint_url)
88-
if ( # pragma: no cover
89-
url_parsed.netloc != endpoint_parsed.netloc
90-
or url_parsed.scheme != endpoint_parsed.scheme
91-
):
92-
error_msg = ( # pragma: no cover
93-
f"Endpoint origin does not match connection origin: {endpoint_url}"
94-
)
95-
logger.error(error_msg) # pragma: no cover
96-
raise ValueError(error_msg) # pragma: no cover
97-
98-
if on_session_created:
99-
session_id = _extract_session_id_from_endpoint(endpoint_url)
100-
if session_id:
101-
on_session_created(session_id)
102-
103-
task_status.started(endpoint_url)
104-
105-
case "message":
106-
# Skip empty data (keep-alive pings)
107-
if not sse.data:
108-
continue
109-
try:
110-
message = types.jsonrpc_message_adapter.validate_json(sse.data, by_name=False)
111-
logger.debug(f"Received server message: {message}")
112-
except Exception as exc: # pragma: no cover
113-
logger.exception("Error parsing server message") # pragma: no cover
114-
await read_stream_writer.send(exc) # pragma: no cover
115-
continue # pragma: no cover
116-
117-
session_message = SessionMessage(message)
118-
await read_stream_writer.send(session_message)
119-
case _: # pragma: no cover
120-
logger.warning(f"Unknown SSE event: {sse.event}") # pragma: no cover
80+
async with sse_events(event_source) as events:
81+
async for sse in events: # pragma: no branch
82+
logger.debug(f"Received SSE event: {sse.event}")
83+
match sse.event:
84+
case "endpoint":
85+
endpoint_url = urljoin(url, sse.data)
86+
logger.debug(f"Received endpoint URL: {endpoint_url}")
87+
88+
url_parsed = urlparse(url)
89+
endpoint_parsed = urlparse(endpoint_url)
90+
if ( # pragma: no cover
91+
url_parsed.netloc != endpoint_parsed.netloc
92+
or url_parsed.scheme != endpoint_parsed.scheme
93+
):
94+
error_msg = ( # pragma: no cover
95+
f"Endpoint origin does not match connection origin: {endpoint_url}"
96+
)
97+
logger.error(error_msg) # pragma: no cover
98+
raise ValueError(error_msg) # pragma: no cover
99+
100+
if on_session_created:
101+
session_id = _extract_session_id_from_endpoint(endpoint_url)
102+
if session_id:
103+
on_session_created(session_id)
104+
105+
task_status.started(endpoint_url)
106+
107+
case "message":
108+
# Skip empty data (keep-alive pings)
109+
if not sse.data:
110+
continue
111+
try:
112+
message = types.jsonrpc_message_adapter.validate_json(sse.data, by_name=False)
113+
logger.debug(f"Received server message: {message}")
114+
except Exception as exc: # pragma: no cover
115+
logger.exception("Error parsing server message") # pragma: no cover
116+
await read_stream_writer.send(exc) # pragma: no cover
117+
continue # pragma: no cover
118+
119+
session_message = SessionMessage(message)
120+
await read_stream_writer.send(session_message)
121+
case _: # pragma: no cover
122+
logger.warning(f"Unknown SSE event: {sse.event}") # pragma: no cover
121123
except SSEError as sse_exc: # pragma: lax no cover
122124
logger.exception("Encountered SSE exception")
123125
raise sse_exc

‎src/mcp/client/streamable_http.py‎

Lines changed: 35 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,7 @@
3737
create_mcp_http_client,
3838
redirect_location,
3939
request_within_origin,
40+
sse_events,
4041
sse_within_origin,
4142
stream_within_origin,
4243
)
@@ -231,15 +232,18 @@ async def handle_get_stream(self, client: httpx2.AsyncClient, read_stream_writer
231232
if last_event_id:
232233
headers[LAST_EVENT_ID] = last_event_id
233234

234-
async with sse_within_origin(client, self.url, headers=headers) as event_source:
235+
async with (
236+
sse_within_origin(client, self.url, headers=headers) as event_source,
237+
sse_events(event_source) as events,
238+
):
235239
if (redirect := _unfollowed_redirect(event_source.response)) is not None:
236240
# The same GET would be redirected again, so retrying cannot help.
237241
logger.warning(f"GET stream not opened: {redirect}")
238242
return
239243
event_source.response.raise_for_status()
240244
logger.debug("GET SSE connection established")
241245

242-
async for sse in event_source:
246+
async for sse in events:
243247
# Track last event ID for reconnection
244248
if sse.id:
245249
last_event_id = sse.id
@@ -278,7 +282,10 @@ async def _handle_resumption_request(self, ctx: RequestContext) -> None:
278282
if isinstance(ctx.session_message.message, JSONRPCRequest): # pragma: no branch
279283
original_request_id = ctx.session_message.message.id
280284

281-
async with sse_within_origin(ctx.client, self.url, headers=headers) as event_source:
285+
async with (
286+
sse_within_origin(ctx.client, self.url, headers=headers) as event_source,
287+
sse_events(event_source) as events,
288+
):
282289
if (redirect := _unfollowed_redirect(event_source.response)) is not None:
283290
logger.warning(redirect)
284291
assert original_request_id is not None
@@ -289,7 +296,7 @@ async def _handle_resumption_request(self, ctx: RequestContext) -> None:
289296
event_source.response.raise_for_status()
290297
logger.debug("Resumption GET SSE connection established")
291298

292-
async for sse in event_source: # pragma: no branch
299+
async for sse in events: # pragma: no branch
293300
is_complete = await self._handle_sse_event(
294301
sse,
295302
ctx.read_stream_writer,
@@ -464,27 +471,27 @@ async def _handle_sse_response(
464471
original_request_id = ctx.session_message.message.id
465472

466473
try:
467-
event_source = EventSource(response)
468-
async for sse in event_source: # pragma: no branch
469-
# Track last event ID for potential reconnection
470-
if sse.id:
471-
last_event_id = sse.id
474+
async with sse_events(EventSource(response)) as events:
475+
async for sse in events: # pragma: no branch
476+
# Track last event ID for potential reconnection
477+
if sse.id:
478+
last_event_id = sse.id
472479

473-
# Track retry interval from server
474-
if sse.retry is not None:
475-
retry_interval_ms = sse.retry
480+
# Track retry interval from server
481+
if sse.retry is not None:
482+
retry_interval_ms = sse.retry
476483

477-
is_complete = await self._handle_sse_event(
478-
sse,
479-
ctx.read_stream_writer,
480-
original_request_id=original_request_id,
481-
resumption_callback=(ctx.metadata.on_resumption_token_update if ctx.metadata else None),
482-
)
483-
# If the SSE event indicates completion, like returning response/error
484-
# break the loop
485-
if is_complete:
486-
await response.aclose()
487-
return # Normal completion, no reconnect needed
484+
is_complete = await self._handle_sse_event(
485+
sse,
486+
ctx.read_stream_writer,
487+
original_request_id=original_request_id,
488+
resumption_callback=(ctx.metadata.on_resumption_token_update if ctx.metadata else None),
489+
)
490+
# If the SSE event indicates completion, like returning response/error
491+
# break the loop
492+
if is_complete:
493+
await response.aclose()
494+
return # Normal completion, no reconnect needed
488495
except Exception:
489496
logger.debug("SSE stream ended", exc_info=True) # pragma: lax no cover
490497

@@ -542,15 +549,18 @@ async def _handle_reconnection(
542549
headers[LAST_EVENT_ID] = last_event_id
543550

544551
try:
545-
async with sse_within_origin(ctx.client, self.url, headers=headers) as event_source:
552+
async with (
553+
sse_within_origin(ctx.client, self.url, headers=headers) as event_source,
554+
sse_events(event_source) as events,
555+
):
546556
event_source.response.raise_for_status()
547557
logger.info("Reconnected to SSE stream")
548558

549559
# Track for potential further reconnection
550560
reconnect_last_event_id: str = last_event_id
551561
reconnect_retry_ms = retry_interval_ms
552562

553-
async for sse in event_source:
563+
async for sse in events:
554564
if sse.id: # pragma: no branch
555565
reconnect_last_event_id = sse.id
556566
if sse.retry is not None:

‎src/mcp/shared/_httpx_utils.py‎

Lines changed: 19 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,12 @@
11
"""Utilities for creating and using httpx2 AsyncClient instances in the MCP transports."""
22

33
from abc import ABC, abstractmethod
4-
from collections.abc import AsyncGenerator
4+
from collections.abc import AsyncGenerator, AsyncIterator
55
from contextlib import asynccontextmanager
6-
from typing import Any, Protocol
6+
from typing import Any
77

88
import httpx2
9+
from typing_extensions import Protocol, runtime_checkable
910

1011
__all__ = ["create_mcp_http_client", "MCP_DEFAULT_TIMEOUT", "MCP_DEFAULT_SSE_READ_TIMEOUT"]
1112

@@ -165,6 +166,22 @@ async def sse_within_origin(
165166
yield httpx2.EventSource(response)
166167

167168

169+
@runtime_checkable
170+
class _AsyncClosable(Protocol):
171+
async def aclose(self) -> None: ...
172+
173+
174+
@asynccontextmanager
175+
async def sse_events(source: httpx2.EventSource) -> AsyncGenerator[AsyncIterator[httpx2.ServerSentEvent]]:
176+
"""Close the outer EventSource iterator if supported; HTTPX2 owns its nested iterators."""
177+
events = source.__aiter__()
178+
try:
179+
yield events
180+
finally:
181+
if isinstance(events, _AsyncClosable):
182+
await events.aclose()
183+
184+
168185
def redirect_location(response: httpx2.Response) -> httpx2.URL | None:
169186
"""Where `response` redirects to, for use in a message: without userinfo, query or fragment,
170187
which can carry state that does not belong in an error or a log line. None if not a redirect."""

‎tests/client/test_streamable_http.py‎

Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -55,6 +55,11 @@
5555
from tests.shared.test_dispatcher import Recorder, echo_handlers
5656

5757

58+
@pytest.fixture(autouse=True)
59+
def _module_runner_lease() -> None:
60+
"""Opt out of the shared runner because iterator cleanup parametrizes `anyio_backend`."""
61+
62+
5863
@pytest.mark.parametrize(
5964
("raw", "expected", "wrapped"),
6065
[
@@ -917,6 +922,80 @@ def handler(request: httpx2.Request) -> httpx2.Response:
917922
assert seen == [("GET http://test/mcp", "evt-41")]
918923

919924

925+
@pytest.mark.anyio
926+
@pytest.mark.parametrize("iterator_kind", ["generator", "closable", "plain"])
927+
@pytest.mark.parametrize("anyio_backend", ["asyncio", "trio"])
928+
async def test_resumed_response_accepts_async_iterators_and_closes_them_when_supported(
929+
monkeypatch: pytest.MonkeyPatch, iterator_kind: str
930+
) -> None:
931+
"""SDK-defined: resumption accepts any EventSource async iterator and closes it when supported.
932+
933+
Substitute the public iterator boundary to isolate representation from HTTPX2's nested-generator cleanup.
934+
"""
935+
expected = JSONRPCResponse(jsonrpc="2.0", id="resume-1", result={"ok": True})
936+
event = httpx2.ServerSentEvent(data=expected.model_dump_json(by_alias=True))
937+
body = f"data: {event.data}\n\n"
938+
closed: list[bool] = []
939+
940+
async def generate() -> AsyncIterator[httpx2.ServerSentEvent]:
941+
try:
942+
yield event
943+
finally:
944+
closed.append(True)
945+
946+
class EventIterator:
947+
def __aiter__(self) -> AsyncIterator[httpx2.ServerSentEvent]:
948+
return self
949+
950+
async def __anext__(self) -> httpx2.ServerSentEvent:
951+
return event
952+
953+
class ClosingEventIterator:
954+
def __aiter__(self) -> AsyncIterator[httpx2.ServerSentEvent]:
955+
return self
956+
957+
async def __anext__(self) -> httpx2.ServerSentEvent:
958+
return event
959+
960+
async def aclose(self) -> None:
961+
closed.append(True)
962+
963+
iterators: dict[str, AsyncIterator[httpx2.ServerSentEvent]] = {
964+
"generator": generate(),
965+
"closable": ClosingEventIterator(),
966+
"plain": EventIterator(),
967+
}
968+
969+
def iterate(source: httpx2.EventSource) -> AsyncIterator[httpx2.ServerSentEvent]:
970+
assert source.response.text == body
971+
return iterators[iterator_kind]
972+
973+
monkeypatch.setattr(httpx2.EventSource, "__aiter__", iterate)
974+
token = "evt-41"
975+
976+
def handler(request: httpx2.Request) -> httpx2.Response:
977+
assert request.method == "GET"
978+
assert request.headers["last-event-id"] == token
979+
return httpx2.Response(200, headers={"content-type": "text/event-stream"}, text=body)
980+
981+
with anyio.fail_after(5):
982+
async with (
983+
httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http,
984+
streamable_http_client("http://test/mcp", http_client=http) as (read, write),
985+
):
986+
await write.send(
987+
SessionMessage(
988+
JSONRPCRequest(jsonrpc="2.0", id=expected.id, method="tools/call", params={}),
989+
metadata=ClientMessageMetadata(resumption_token=token),
990+
)
991+
)
992+
reply = await read.receive()
993+
994+
assert isinstance(reply, SessionMessage)
995+
assert reply.message == expected
996+
assert closed == ([] if iterator_kind == "plain" else [True])
997+
998+
920999
async def _redirected_call_error(url: str, location: str) -> str:
9211000
"""Send one request through streamable_http_client to a server answering `url` with a 307 to
9221001
`location`, and return the message of the error that resolves it."""

‎tests/interaction/transports/test_hosting_resume.py‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
"""
1111

1212
import json
13+
from collections.abc import AsyncGenerator
1314

1415
import anyio
1516
import httpx2
@@ -70,9 +71,13 @@ def _tools_call(request_id: int, name: str, arguments: dict[str, object]) -> str
7071

7172

7273
async def _read_events(response: httpx2.Response, count: int) -> list[ServerSentEvent]:
73-
"""Read exactly `count` SSE events from a streaming response without closing it."""
74+
"""Read exactly `count` SSE events and close the iterator."""
7475
source = aiter(EventSource(response))
75-
return [await anext(source) for _ in range(count)]
76+
try:
77+
return [await anext(source) for _ in range(count)]
78+
finally:
79+
assert isinstance(source, AsyncGenerator)
80+
await source.aclose()
7681

7782

7883
@requirement("hosting:resume:event-ids")

‎tests/shared/test_sse.py‎

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -120,8 +120,12 @@ async def test_raw_sse_connection() -> None:
120120
assert response.headers["content-type"] == "text/event-stream; charset=utf-8"
121121

122122
lines = response.aiter_lines()
123-
assert await anext(lines) == "event: endpoint"
124-
assert (await anext(lines)).startswith("data: /messages/?session_id=")
123+
try:
124+
assert await anext(lines) == "event: endpoint"
125+
assert (await anext(lines)).startswith("data: /messages/?session_id=")
126+
finally:
127+
assert isinstance(lines, AsyncGenerator)
128+
await lines.aclose()
125129

126130

127131
@pytest.mark.anyio

0 commit comments

Comments
 (0)