|
1 | 1 | from collections.abc import AsyncGenerator |
2 | | -from typing import Any |
| 2 | +from typing import Any, cast |
3 | 3 |
|
4 | 4 | import anyio |
5 | 5 | import pytest |
|
10 | 10 | from mcp.shared.exceptions import McpError |
11 | 11 | from mcp.shared.memory import create_client_server_memory_streams, create_connected_server_and_client_session |
12 | 12 | from mcp.shared.message import SessionMessage |
| 13 | +from mcp.shared.session import RequestResponder |
13 | 14 | from mcp.types import ( |
14 | 15 | CancelledNotification, |
15 | 16 | CancelledNotificationParams, |
|
25 | 26 | ) |
26 | 27 |
|
27 | 28 |
|
| 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 | + |
28 | 41 | @pytest.fixture |
29 | 42 | def mcp_server() -> Server: |
30 | 43 | return Server(name="test server") |
@@ -128,6 +141,44 @@ async def make_request(client_session: ClientSession): |
128 | 141 | await ev_cancelled.wait() |
129 | 142 |
|
130 | 143 |
|
| 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 | + |
131 | 182 | @pytest.mark.anyio |
132 | 183 | async def test_response_id_type_mismatch_string_to_int(): |
133 | 184 | """ |
|
0 commit comments