Skip to content

Commit 76c33e4

Browse files
author
Jianke LIN
committed
fix(stdio): bound EOF drain wait
1 parent a97b9e5 commit 76c33e4

4 files changed

Lines changed: 104 additions & 18 deletions

File tree

src/mcp/server/lowlevel/server.py

Lines changed: 16 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,8 @@ async def main():
7272

7373
logger = logging.getLogger(__name__)
7474

75+
DEFAULT_READ_EOF_DRAIN_TIMEOUT_SECONDS = 1.0
76+
7577
LifespanResultT = TypeVar("LifespanResultT", default=Any)
7678

7779
_ParamsT = TypeVar("_ParamsT", bound=BaseModel, default=BaseModel)
@@ -695,6 +697,9 @@ async def run(
695697
# to drain their responses via the still-open write stream (e.g. stdio
696698
# with bash-redirected stdin).
697699
drain_on_read_close: bool = False,
700+
# Maximum time to wait for in-flight handlers to drain after read EOF.
701+
# None means wait indefinitely.
702+
read_eof_drain_timeout_seconds: float | None = DEFAULT_READ_EOF_DRAIN_TIMEOUT_SECONDS,
698703
) -> None:
699704
"""Serve a single connection over the given streams until the read side closes.
700705
@@ -705,20 +710,17 @@ async def run(
705710
streamable-HTTP manager) call `serve_loop` directly instead.
706711
"""
707712
async with self.lifespan(self) as lifespan_context:
708-
try:
709-
await serve_dual_era_loop(
710-
self,
711-
read_stream,
712-
write_stream,
713-
lifespan_state=lifespan_context,
714-
init_options=initialization_options,
715-
raise_exceptions=raise_exceptions,
716-
session_id=None,
717-
close_write_stream_on_read_close=not drain_on_read_close,
718-
)
719-
finally:
720-
if drain_on_read_close:
721-
await write_stream.aclose()
713+
await serve_dual_era_loop(
714+
self,
715+
read_stream,
716+
write_stream,
717+
lifespan_state=lifespan_context,
718+
init_options=initialization_options,
719+
raise_exceptions=raise_exceptions,
720+
session_id=None,
721+
close_write_stream_on_read_close=not drain_on_read_close,
722+
read_eof_drain_timeout_seconds=read_eof_drain_timeout_seconds,
723+
)
722724

723725
def streamable_http_app(
724726
self,

src/mcp/server/runner.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -415,6 +415,7 @@ async def serve_loop(
415415
init_options: InitializationOptions | None = None,
416416
raise_exceptions: bool = False,
417417
close_write_stream_on_read_close: bool = True,
418+
read_eof_drain_timeout_seconds: float | None = None,
418419
) -> None:
419420
"""Drive ``server`` in handshake-only loop mode over a stream pair until the channel closes.
420421
@@ -434,6 +435,7 @@ async def serve_loop(
434435
# state instead of failing the init-gate.
435436
inline_methods=frozenset({"initialize"}),
436437
close_write_stream_on_read_close=close_write_stream_on_read_close,
438+
read_eof_drain_timeout_seconds=read_eof_drain_timeout_seconds,
437439
)
438440
connection = Connection.for_loop(dispatcher, session_id=session_id)
439441
await serve_connection(
@@ -548,6 +550,7 @@ async def serve_dual_era_loop(
548550
init_options: InitializationOptions | None = None,
549551
raise_exceptions: bool = False,
550552
close_write_stream_on_read_close: bool = True,
553+
read_eof_drain_timeout_seconds: float | None = None,
551554
) -> None:
552555
"""Drive `server` over a duplex stream pair, serving both protocol eras.
553556
@@ -597,6 +600,7 @@ async def serve_dual_era_loop(
597600
# next pipelined message is read.
598601
inline_methods=frozenset({"initialize", "server/discover"}),
599602
close_write_stream_on_read_close=close_write_stream_on_read_close,
603+
read_eof_drain_timeout_seconds=read_eof_drain_timeout_seconds,
600604
)
601605
loop_connection = Connection.for_loop(dispatcher, session_id=session_id)
602606
loop_runner = ServerRunner(server, loop_connection, lifespan_state, init_options=init_options)

src/mcp/shared/jsonrpc_dispatcher.py

Lines changed: 13 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -252,6 +252,7 @@ def __init__(
252252
raise_handler_exceptions: bool = False,
253253
inline_methods: frozenset[str] = frozenset(),
254254
close_write_stream_on_read_close: bool = True,
255+
read_eof_drain_timeout_seconds: float | None = None,
255256
on_stream_exception: Callable[[Exception], Awaitable[None]] | None = None,
256257
) -> None:
257258
"""Wire a dispatcher over a transport's `SessionMessage` stream pair.
@@ -268,6 +269,9 @@ def __init__(
268269
read stream closes. Full-duplex transports may set this to
269270
false so in-flight handlers can finish writing responses after
270271
input EOF.
272+
read_eof_drain_timeout_seconds: Maximum time to wait for in-flight
273+
handlers to drain after read EOF when the write stream stays
274+
open; None waits indefinitely.
271275
on_stream_exception: Observer for `Exception` items on the read
272276
stream; without it they are debug-logged and dropped. Awaited
273277
inline in the read loop, so a slow observer stalls dispatch.
@@ -283,6 +287,7 @@ def __init__(
283287
self._peer_cancel_mode: PeerCancelMode = peer_cancel_mode
284288
self._raise_handler_exceptions = raise_handler_exceptions
285289
self._close_write_stream_on_read_close = close_write_stream_on_read_close
290+
self._read_eof_drain_timeout_seconds = read_eof_drain_timeout_seconds
286291
self._inline_methods = inline_methods
287292
self.on_stream_exception = on_stream_exception
288293
"""Observer for ``Exception`` items on the read stream. Mutable so a session can
@@ -481,16 +486,21 @@ async def run(
481486
self._fan_out_closed()
482487
normal_eof = True
483488
finally:
484-
if not normal_eof:
485-
# Cancel on crash/cancel paths. On normal EOF, let
486-
# already received handlers drain their responses.
489+
if not normal_eof or self._close_write_stream_on_read_close:
490+
# Cancel on crash/cancel paths. If read EOF also closed
491+
# writes, handlers cannot drain responses anyway.
487492
tg.cancel_scope.cancel()
493+
elif self._read_eof_drain_timeout_seconds is not None:
494+
tg.cancel_scope.deadline = anyio.current_time() + self._read_eof_drain_timeout_seconds
488495
finally:
489496
# Covers cancel/crash paths that skip the inline fan-out; idempotent.
490497
self._running = False
491498
self._closed = True
492499
self._tg = None
493500
self._fan_out_closed()
501+
if not self._close_write_stream_on_read_close:
502+
with anyio.CancelScope(shield=True):
503+
await self._write_stream.aclose()
494504
await resync_tracer()
495505

496506
async def _dispatch(

tests/server/test_cancel_handling.py

Lines changed: 71 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -120,7 +120,13 @@ async def handle_call_tool(ctx: ServerRequestContext, params: CallToolRequestPar
120120
server_write, from_server = anyio.create_memory_object_stream[SessionMessage](10)
121121

122122
async def run_server():
123-
await server.run(server_read, server_write, server.create_initialization_options(), drain_on_read_close=True)
123+
await server.run(
124+
server_read,
125+
server_write,
126+
server.create_initialization_options(),
127+
drain_on_read_close=True,
128+
read_eof_drain_timeout_seconds=None,
129+
)
124130
server_run_returned.set()
125131

126132
init_req = JSONRPCRequest(
@@ -166,6 +172,70 @@ async def run_server():
166172
await server_run_returned.wait()
167173

168174

175+
@pytest.mark.anyio
176+
async def test_server_bounds_drain_on_read_eof_when_handler_never_finishes():
177+
handler_started = anyio.Event()
178+
handler_cancelled = anyio.Event()
179+
server_run_returned = anyio.Event()
180+
181+
async def handle_call_tool(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult:
182+
handler_started.set()
183+
try:
184+
await anyio.sleep_forever()
185+
finally:
186+
handler_cancelled.set()
187+
raise AssertionError # pragma: no cover
188+
189+
server = Server("test", on_call_tool=handle_call_tool)
190+
191+
to_server, server_read = anyio.create_memory_object_stream[SessionMessage | Exception](10)
192+
server_write, from_server = anyio.create_memory_object_stream[SessionMessage](10)
193+
194+
async def run_server():
195+
await server.run(
196+
server_read,
197+
server_write,
198+
server.create_initialization_options(),
199+
drain_on_read_close=True,
200+
read_eof_drain_timeout_seconds=0.05,
201+
)
202+
server_run_returned.set()
203+
204+
init_req = JSONRPCRequest(
205+
jsonrpc="2.0",
206+
id=1,
207+
method="initialize",
208+
params=InitializeRequestParams(
209+
protocol_version=LATEST_HANDSHAKE_VERSION,
210+
capabilities=ClientCapabilities(),
211+
client_info=Implementation(name="test", version="1.0"),
212+
).model_dump(by_alias=True, mode="json", exclude_none=True),
213+
)
214+
initialized = JSONRPCNotification(jsonrpc="2.0", method="notifications/initialized")
215+
call_req = JSONRPCRequest(
216+
jsonrpc="2.0",
217+
id=2,
218+
method="tools/call",
219+
params=CallToolRequestParams(name="slow", arguments={}).model_dump(by_alias=True, mode="json"),
220+
)
221+
222+
with anyio.fail_after(2):
223+
async with anyio.create_task_group() as tg, to_server, server_read, server_write, from_server:
224+
tg.start_soon(run_server)
225+
226+
await to_server.send(SessionMessage(init_req))
227+
await from_server.receive() # init response
228+
await to_server.send(SessionMessage(initialized))
229+
await to_server.send(SessionMessage(call_req))
230+
231+
await handler_started.wait()
232+
await to_server.aclose()
233+
234+
await server_run_returned.wait()
235+
236+
assert handler_cancelled.is_set()
237+
238+
169239
@pytest.mark.anyio
170240
async def test_server_reraises_handler_cancellation_when_server_is_cancelled():
171241
"""If the server task is cancelled (e.g. KeyboardInterrupt), in-flight

0 commit comments

Comments
 (0)