Skip to content

Commit 6922ff5

Browse files
committed
fix(direct-dispatcher): contain callback failures
1 parent 9972c21 commit 6922ff5

2 files changed

Lines changed: 56 additions & 2 deletions

File tree

src/mcp/shared/direct_dispatcher.py

Lines changed: 27 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
from mcp.shared._compat import resync_tracer
3131
from mcp.shared.dispatcher import (
3232
CallOptions,
33+
DispatchContext,
3334
OnNotify,
3435
OnNotifyIntercept,
3536
OnRequest,
@@ -43,6 +44,30 @@
4344

4445
logger = logging.getLogger(__name__)
4546

47+
48+
def _shielded_progress(fn: ProgressFnT) -> ProgressFnT:
49+
"""Wrap a progress callback so its failure does not fail the request."""
50+
51+
async def _wrapped(progress: float, total: float | None, message: str | None) -> None:
52+
try:
53+
await fn(progress, total, message)
54+
except Exception:
55+
logger.exception("progress callback raised")
56+
57+
return _wrapped
58+
59+
60+
def _contained_notify(fn: OnNotify) -> OnNotify:
61+
"""Wrap a notification handler so its failure does not reach the sender."""
62+
63+
async def _wrapped(dctx: DispatchContext[TransportContext], method: str, params: Mapping[str, Any] | None) -> None:
64+
try:
65+
await fn(dctx, method, params)
66+
except Exception:
67+
logger.exception("notification handler for %r raised", method)
68+
69+
return _wrapped
70+
4671
__all__ = ["DirectDispatcher", "create_direct_dispatcher_pair"]
4772

4873
DIRECT_TRANSPORT_KIND = "direct"
@@ -206,7 +231,7 @@ def _make_context(
206231
_back_request=lambda m, p, o: peer._dispatch_request(m, p, o),
207232
_back_notify=lambda m, p: peer._dispatch_notify(m, p),
208233
request_id=request_id,
209-
_on_progress=on_progress,
234+
_on_progress=_shielded_progress(on_progress) if on_progress is not None else None,
210235
)
211236

212237
async def _wait_ready(self) -> None:
@@ -301,7 +326,7 @@ async def _dispatch_notify(self, method: str, params: Mapping[str, Any] | None)
301326
return
302327
assert self._on_notify is not None
303328
dctx = self._make_context()
304-
await self._on_notify(dctx, method, params)
329+
await _contained_notify(self._on_notify)(dctx, method, params)
305330

306331

307332
def create_direct_dispatcher_pair(

tests/shared/test_dispatcher.py

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -216,6 +216,35 @@ async def on_progress(progress: float, total: float | None, message: str | None)
216216
assert received == [(0.5, 1.0, "halfway")]
217217

218218

219+
@pytest.mark.anyio
220+
async def test_progress_callback_exception_does_not_fail_request(pair_factory: PairFactory):
221+
async def server_on_request(
222+
ctx: DispatchContext[TransportContext], method: str, params: Mapping[str, Any] | None
223+
) -> dict[str, Any]:
224+
await ctx.progress(0.5)
225+
return {"ok": True}
226+
227+
async def on_progress(progress: float, total: float | None, message: str | None) -> None:
228+
raise RuntimeError("consumer failed")
229+
230+
async with running_pair(pair_factory, server_on_request=server_on_request) as (client, *_):
231+
with anyio.fail_after(5):
232+
result = await client.send_raw_request("tools/call", None, {"on_progress": on_progress})
233+
assert result == {"ok": True}
234+
235+
236+
@pytest.mark.anyio
237+
async def test_notification_handler_exception_does_not_reach_sender(pair_factory: PairFactory):
238+
async def on_notify(
239+
ctx: DispatchContext[TransportContext], method: str, params: Mapping[str, Any] | None
240+
) -> None:
241+
raise RuntimeError("handler failed")
242+
243+
async with running_pair(pair_factory, client_on_notify=on_notify) as (client, *_):
244+
with anyio.fail_after(5):
245+
await client.notify("notifications/message", None)
246+
247+
219248
@pytest.mark.anyio
220249
async def test_ctx_progress_is_noop_when_caller_supplied_no_callback(pair_factory: PairFactory):
221250
async def server_on_request(

0 commit comments

Comments
 (0)