Skip to content

Commit 6cedc87

Browse files
committed
Fix completed request cancellation cleanup
1 parent 98b7159 commit 6cedc87

2 files changed

Lines changed: 57 additions & 2 deletions

File tree

src/mcp/shared/session.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -114,7 +114,11 @@ def __exit__(
114114
self._entered = False
115115
if not self._cancel_scope: # pragma: no cover
116116
raise RuntimeError("No active cancel scope")
117-
self._cancel_scope.__exit__(exc_type, exc_val, exc_tb)
117+
try:
118+
self._cancel_scope.__exit__(exc_type, exc_val, exc_tb)
119+
except BaseException as exc:
120+
if not (self._completed and isinstance(exc, anyio.get_cancelled_exc_class())):
121+
raise
118122

119123
async def respond(self, response: SendResultT | ErrorData) -> None:
120124
"""Send a response for this request.

tests/shared/test_session.py

Lines changed: 52 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
from collections.abc import AsyncGenerator
2-
from typing import Any
2+
from typing import Any, cast
33

44
import anyio
55
import pytest
@@ -10,6 +10,7 @@
1010
from mcp.shared.exceptions import McpError
1111
from mcp.shared.memory import create_client_server_memory_streams, create_connected_server_and_client_session
1212
from mcp.shared.message import SessionMessage
13+
from mcp.shared.session import RequestResponder
1314
from mcp.types import (
1415
CancelledNotification,
1516
CancelledNotificationParams,
@@ -25,6 +26,18 @@
2526
)
2627

2728

29+
class _CancelScopeThatRaisesOnExit:
30+
cancel_called = True
31+
32+
def __exit__(
33+
self,
34+
exc_type: type[BaseException] | None,
35+
exc_val: BaseException | None,
36+
exc_tb: object | None,
37+
) -> None:
38+
raise anyio.get_cancelled_exc_class()()
39+
40+
2841
@pytest.fixture
2942
def mcp_server() -> Server:
3043
return Server(name="test server")
@@ -128,6 +141,44 @@ async def make_request(client_session: ClientSession):
128141
await ev_cancelled.wait()
129142

130143

144+
@pytest.mark.anyio
145+
async def test_completed_request_responder_suppresses_cancel_scope_exit() -> None:
146+
completed: list[Any] = []
147+
responder = RequestResponder(
148+
request_id=1,
149+
request_meta=None,
150+
request=types.ClientRequest(types.PingRequest()),
151+
session=cast(Any, object()),
152+
on_complete=completed.append,
153+
)
154+
responder._completed = True # type: ignore[reportPrivateUsage]
155+
responder._cancel_scope = cast( # type: ignore[reportPrivateUsage]
156+
anyio.CancelScope, _CancelScopeThatRaisesOnExit()
157+
)
158+
159+
responder.__exit__(None, None, None)
160+
161+
assert completed == [responder]
162+
assert not responder._entered # type: ignore[reportPrivateUsage]
163+
164+
165+
@pytest.mark.anyio
166+
async def test_incomplete_request_responder_propagates_cancel_scope_exit() -> None:
167+
responder = RequestResponder(
168+
request_id=1,
169+
request_meta=None,
170+
request=types.ClientRequest(types.PingRequest()),
171+
session=cast(Any, object()),
172+
on_complete=lambda _: None,
173+
)
174+
responder._cancel_scope = cast( # type: ignore[reportPrivateUsage]
175+
anyio.CancelScope, _CancelScopeThatRaisesOnExit()
176+
)
177+
178+
with pytest.raises(anyio.get_cancelled_exc_class()):
179+
responder.__exit__(None, None, None)
180+
181+
131182
@pytest.mark.anyio
132183
async def test_response_id_type_mismatch_string_to_int():
133184
"""

0 commit comments

Comments
 (0)