diff --git a/docs/experimental-v2.md b/docs/experimental-v2.md index 7c1149a..e62ecb6 100644 --- a/docs/experimental-v2.md +++ b/docs/experimental-v2.md @@ -150,6 +150,8 @@ Only the initial v2 request is reduced to the common v1 initialization fields when an agent selects v1. Client-side fallback is application controlled and may require opening a new -transport. Protocol-level request cancellation is not yet exposed by the -experimental runtime; `session/cancel` remains available for cancelling active -session work. +transport. Protocol-level request cancellation (`$/cancel_request`) is handled +by the shared connection layer, as in v1: cancelling the task awaiting a request +sends a best-effort cancellation to the peer, and an incoming cancellation +cancels the handler task, which answers with its result or `-32800`. +`session/cancel` remains available for cancelling active session work. diff --git a/docs/quickstart.md b/docs/quickstart.md index 15cfb05..14c43d9 100644 --- a/docs/quickstart.md +++ b/docs/quickstart.md @@ -228,6 +228,21 @@ MCP requests return the inner JSON result unchanged, including `null`. Use and optional `params`. These methods share the same connections and routers across stdio, HTTP, and WebSocket transports. +## Request cancellation + +Connections handle the protocol's `$/cancel_request` notification. When the +peer cancels one of its requests, the SDK cancels the task running your handler. +The handler may catch `asyncio.CancelledError` and return a (partial) result; +otherwise the peer receives a `-32800` "Request cancelled" error. Cancellations +for unknown or already finished requests are ignored; a cancelled request still +gets exactly one response unless the connection is closed first. + +Cancelling the task that awaits an outgoing request (for example through +`asyncio.wait_for`) still raises `CancelledError` locally without waiting for the +peer, and additionally sends a best-effort `$/cancel_request` for that request. +Cancellation support is optional for peers. Use `session/cancel` to stop a +prompt turn. + ## Maintaining protocol routes The `Agent` and `Client` protocols in `src/acp/interfaces.py` are the source of diff --git a/src/acp/connection.py b/src/acp/connection.py index 84f8f05..a827210 100644 --- a/src/acp/connection.py +++ b/src/acp/connection.py @@ -14,6 +14,7 @@ from ._transport import NdjsonTransport, Transport from .exceptions import RequestError +from .meta import PROTOCOL_METHODS from .task import MessageSender, TaskSupervisor from .telemetry import span_context @@ -37,6 +38,13 @@ class StreamEvent: StreamObserver = Callable[[StreamEvent], Awaitable[None] | None] +_CANCEL_REQUEST_METHOD = PROTOCOL_METHODS["cancel_request"] + + +def _is_request_id(value: Any) -> bool: + """Whether ``value`` is a JSON-RPC ``RequestId``: ``null``, an integer or a string.""" + return value is None or isinstance(value, str) or (isinstance(value, int) and not isinstance(value, bool)) + class Connection: """Minimal JSON-RPC 2.0 connection over newline-delimited JSON frames.""" @@ -54,6 +62,7 @@ def __init__( self._handler = handler self._next_request_id = 0 self._pending: dict[int, asyncio.Future[Any]] = {} + self._incoming: dict[Any, asyncio.Task[Any]] = {} self._tasks = TaskSupervisor(source="acp.Connection") self._tasks.add_error_handler(self._on_task_error) self._closed = False @@ -117,16 +126,14 @@ async def send_request(self, method: str, params: JsonValue | None = None) -> An payload = {"jsonrpc": "2.0", "id": request_id, "method": method, "params": params} try: await self._transport.send(payload) - except BaseException: - self._pending.pop(request_id, None) - future.cancel() - raise - self._notify_observers(StreamDirection.OUTGOING, payload) - try: + self._notify_observers(StreamDirection.OUTGOING, payload) return await future - except asyncio.CancelledError: - self._pending.pop(request_id, None) + except BaseException as exc: + still_pending = self._pending.pop(request_id, None) is not None future.cancel() + if isinstance(exc, asyncio.CancelledError) and still_pending: + # The request may already be on the wire; ask the peer to stop working on it. + self._send_cancel_request(request_id) raise async def send_notification(self, method: str, params: JsonValue | None = None) -> None: @@ -152,13 +159,20 @@ async def _receive_loop(self) -> None: def _process_message(self, message: dict[str, Any]) -> None: method = message.get("method") has_id = "id" in message + if method == _CANCEL_REQUEST_METHOD and not has_id: + self._cancel_incoming(message.get("params")) + return if method is not None: # this is a request or notification # {"jsonrpc": "2.0", "id": 1, "method": "foo", "params": {...}} # request # {"jsonrpc": "2.0", "method": "foo", "params: {...}} # notification - self._tasks.create( - self._run_request(message) if has_id else self._run_notification(message), - name="acp.Connection.request" if has_id else "acp.Connection.notification", - ) + if not has_id: + self._tasks.create(self._run_notification(message), name="acp.Connection.notification") + return + # The handler gets its own task so ``$/cancel_request`` can cancel it, even before it + # starts, without also cancelling delivery of the response the peer still expects. + handler = self._tasks.create(self._execute_request(message), name="acp.Connection.request") + self._track_incoming(message["id"], handler) + self._tasks.create(self._run_request(message, handler), name="acp.Connection.response") return if has_id: # this is a response, {"id", "result" | "error"} self._handle_response(message) @@ -184,8 +198,56 @@ def _notify_observers(self, direction: StreamDirection, message: dict[str, Any]) def _on_observer_error(self, task: asyncio.Task[Any], exc: BaseException) -> None: logging.exception("Stream observer coroutine failed", exc_info=exc) - async def _run_request(self, message: dict[str, Any]) -> None: - payload = await self._execute_request(message) + def _track_incoming(self, request_id: Any, task: asyncio.Task[Any]) -> None: + # A ``$/cancel_request`` can only name a valid request id, and Python would alias an invalid + # one such as ``true`` or ``1.0`` with the integer ``1``. + if not _is_request_id(request_id): + return + self._incoming[request_id] = task + + def _forget(done: asyncio.Task[Any]) -> None: + if self._incoming.get(request_id) is done: + del self._incoming[request_id] + + task.add_done_callback(_forget) + + def _cancel_incoming(self, params: Any) -> None: + # ``requestId`` is required and may be ``null``, so a missing one must not match a ``null`` id. + if not isinstance(params, dict) or "requestId" not in params: + return + request_id = params["requestId"] + if not _is_request_id(request_id): + return + # Unknown or already finished requests are ignored, as the protocol allows. + task = self._incoming.get(request_id) + if task is not None: + task.cancel() + + def _send_cancel_request(self, request_id: int) -> None: + if self._closed or self._disconnected: + return + self._tasks.create( + self.send_notification(_CANCEL_REQUEST_METHOD, {"requestId": request_id}), + name="acp.Connection.cancel_request", + on_error=self._on_cancel_request_error, + ) + + def _on_cancel_request_error(self, task: asyncio.Task[Any], exc: BaseException) -> None: + logging.debug("Failed to send %s", _CANCEL_REQUEST_METHOD, exc_info=exc) + + async def _run_request( + self, message: dict[str, Any], handler: asyncio.Future[dict[str, Any]] | None = None + ) -> None: + if handler is None: + handler = self._tasks.create(self._execute_request(message), name="acp.Connection.request") + # Unlike ``await handler``, ``asyncio.wait`` does not re-raise the handler's cancellation, and + # ``$/cancel_request`` only targets the handler, so only ``close()`` can abandon the response. + await asyncio.wait({handler}) + if handler.cancelled(): + # Cancelled by ``$/cancel_request`` or from inside the handler. + payload = {"jsonrpc": "2.0", "id": message["id"], "error": RequestError.request_cancelled().to_error_obj()} + else: + payload = handler.result() await self._transport.send(payload) self._notify_observers(StreamDirection.OUTGOING, payload) diff --git a/src/acp/exceptions.py b/src/acp/exceptions.py index 06098dd..c6f9083 100644 --- a/src/acp/exceptions.py +++ b/src/acp/exceptions.py @@ -42,5 +42,9 @@ def resource_not_found(cls, uri: str | None = None) -> RequestError: data = {"uri": uri} if uri is not None else None return cls(-32002, "Resource not found", data) + @classmethod + def request_cancelled(cls, data: dict[str, Any] | None = None) -> RequestError: + return cls(-32800, "Request cancelled", data) + def to_error_obj(self) -> dict[str, Any]: return {"code": self.code, "message": str(self), "data": self.data} diff --git a/tests/test_request_cancellation.py b/tests/test_request_cancellation.py new file mode 100644 index 0000000..2206e7e --- /dev/null +++ b/tests/test_request_cancellation.py @@ -0,0 +1,459 @@ +"""``$/cancel_request`` handling (https://agentclientprotocol.com/protocol/v1/cancellation).""" + +from __future__ import annotations + +import asyncio +import json +import logging +from typing import Any, cast + +import pytest + +from acp import Agent +from acp.connection import Connection +from acp.core import AgentSideConnection, ClientSideConnection +from acp.schema import PermissionOption, ToolCallUpdate +from tests.conftest import TestAgent, TestClient + + +async def _write(writer: asyncio.StreamWriter, message: dict[str, Any]) -> None: + writer.write((json.dumps(message) + "\n").encode()) + await writer.drain() + + +async def _read(reader: asyncio.StreamReader) -> dict[str, Any]: + return json.loads(await asyncio.wait_for(reader.readline(), timeout=1)) + + +def _prompt(request_id: Any) -> dict[str, Any]: + return { + "jsonrpc": "2.0", + "id": request_id, + "method": "session/prompt", + "params": {"sessionId": "sess", "prompt": [{"type": "text", "text": "hi"}]}, + } + + +def _cancel_request(request_id: Any) -> dict[str, Any]: + return {"jsonrpc": "2.0", "method": "$/cancel_request", "params": {"requestId": request_id}} + + +class _BlockingAgent(TestAgent): + """Blocks every prompt until cancelled and records how each one ended.""" + + def __init__(self) -> None: + super().__init__() + self.started = asyncio.Event() + self.release = asyncio.Event() + self.cancelled: list[str] = [] + + async def prompt(self, session_id: str, prompt: list[Any], **kwargs: Any) -> Any: + self.started.set() + try: + await self.release.wait() + except asyncio.CancelledError: + self.cancelled.append(session_id) + raise + return await super().prompt(session_id, prompt, **kwargs) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_id", [7, "req-7"]) +async def test_cancel_request_cancels_handler_and_replies_request_cancelled( + server, caplog: pytest.LogCaptureFixture, request_id: int | str +) -> None: + agent = _BlockingAgent() + async with AgentSideConnection(cast(Agent, agent), server.server_writer, server.server_reader, listening=True): + with caplog.at_level(logging.ERROR): + await _write(server.client_writer, _prompt(request_id)) + await asyncio.wait_for(agent.started.wait(), timeout=1) + await _write(server.client_writer, _cancel_request(request_id)) + response = await _read(server.client_reader) + + assert response["id"] == request_id + assert response["error"]["code"] == -32800 + assert response["error"]["message"] == "Request cancelled" + assert "result" not in response + assert agent.cancelled == ["sess"] + assert "$/cancel_request" not in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_id", [0, "req-0", None]) +async def test_cancel_request_before_the_handler_starts_still_replies( + server, caplog: pytest.LogCaptureFixture, request_id: int | str | None +) -> None: + agent = _BlockingAgent() + async with AgentSideConnection(cast(Agent, agent), server.server_writer, server.server_reader, listening=True): + with caplog.at_level(logging.ERROR): + # One write, so the receive loop reads both frames before the handler task first runs. + server.client_writer.write( + (json.dumps(_prompt(request_id)) + "\n" + json.dumps(_cancel_request(request_id)) + "\n").encode() + ) + await server.client_writer.drain() + response = await _read(server.client_reader) + + assert response["id"] == request_id + assert response["error"]["code"] == -32800 + assert not agent.started.is_set() + assert caplog.text == "" + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(server.client_reader.readline(), timeout=0.1) + + +class _GatedTransport: + """Message transport whose sends block until ``release`` is set.""" + + def __init__(self) -> None: + self.incoming: asyncio.Queue[dict[str, Any]] = asyncio.Queue() + self.sent: list[dict[str, Any]] = [] + self.sending = asyncio.Event() + self.release = asyncio.Event() + self.settled = asyncio.Event() + self.send_cancelled = False + self._receiving = False + + async def send(self, message: dict[str, Any]) -> None: + self.sending.set() + try: + await self.release.wait() + except asyncio.CancelledError: + self.send_cancelled = True + raise + else: + self.sent.append(message) + finally: + self.settled.set() + + async def receive(self) -> dict[str, Any] | None: + self._receiving = True + try: + return await self.incoming.get() + finally: + self._receiving = False + + async def close(self) -> None: + pass + + async def deliver(self, message: dict[str, Any]) -> None: + """Queue ``message`` and wait until the connection has processed it.""" + await self.incoming.put(message) + for _ in range(100): + if self._receiving and self.incoming.empty(): + return + await asyncio.sleep(0) + raise AssertionError("the connection did not process the message") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("handler_cancelled", [False, True], ids=["result", "request_cancelled"]) +async def test_cancel_request_during_response_send_keeps_the_response(handler_cancelled: bool) -> None: + transport = _GatedTransport() + started = asyncio.Event() + + async def handler(method: str, params: Any, is_notification: bool) -> Any: + started.set() + if handler_cancelled: + await asyncio.Event().wait() + return {"ok": True} + + async with Connection(handler, transport): + await transport.deliver(_prompt(0)) + await asyncio.wait_for(started.wait(), timeout=1) + if handler_cancelled: + await transport.deliver(_cancel_request(0)) + await asyncio.wait_for(transport.sending.wait(), timeout=1) + + # A late (or repeated) cancellation lands while the response is still being sent. + await transport.deliver(_cancel_request(0)) + assert transport.sent == [] + transport.release.set() + await asyncio.wait_for(transport.settled.wait(), timeout=1) + assert not transport.send_cancelled, "the cancellation aborted the response send" + + if handler_cancelled: + assert [(m["id"], m["error"]["code"]) for m in transport.sent] == [(0, -32800)] + else: + assert transport.sent == [{"jsonrpc": "2.0", "id": 0, "result": {"ok": True}}] + + +@pytest.mark.asyncio +async def test_close_cancels_a_blocked_response_send() -> None: + transport = _GatedTransport() + + async def handler(method: str, params: Any, is_notification: bool) -> Any: + return {"ok": True} + + conn = Connection(handler, transport) + await transport.deliver(_prompt(0)) + await asyncio.wait_for(transport.sending.wait(), timeout=1) + + await asyncio.wait_for(conn.close(), timeout=1) + + assert transport.send_cancelled + assert transport.sent == [] + + +@pytest.mark.asyncio +async def test_cancel_request_only_affects_the_targeted_request(server, caplog: pytest.LogCaptureFixture) -> None: + agent = _BlockingAgent() + async with AgentSideConnection(cast(Agent, agent), server.server_writer, server.server_reader, listening=True): + with caplog.at_level(logging.ERROR): + await _write(server.client_writer, _prompt(1)) + await asyncio.wait_for(agent.started.wait(), timeout=1) + # Unknown and already-finished ids are ignored without a reply or an error log. + await _write(server.client_writer, _cancel_request(99)) + await _write(server.client_writer, {"jsonrpc": "2.0", "id": 2, "method": "session/list", "params": {}}) + listed = await _read(server.client_reader) + await _write(server.client_writer, _cancel_request(2)) + + agent.release.set() + finished = await _read(server.client_reader) + + assert listed == {"jsonrpc": "2.0", "id": 2, "result": {"sessions": []}} + assert finished == {"jsonrpc": "2.0", "id": 1, "result": {"stopReason": "end_turn"}} + assert agent.cancelled == [] + assert caplog.text == "" + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(server.client_reader.readline(), timeout=0.1) + + +async def _cancel_while_pending(request_id: Any, cancel: dict[str, Any]) -> dict[str, Any]: + """Deliver ``cancel`` while request ``request_id`` is in flight, then let it finish and return its response.""" + transport = _GatedTransport() + transport.release.set() + started = asyncio.Event() + release = asyncio.Event() + + async def handler(method: str, params: Any, is_notification: bool) -> Any: + started.set() + await release.wait() + return {"ok": True} + + async with Connection(handler, transport): + await transport.deliver(_prompt(request_id)) + await asyncio.wait_for(started.wait(), timeout=1) + await transport.deliver(cancel) + release.set() + await asyncio.wait_for(transport.settled.wait(), timeout=1) + + [response] = transport.sent + return response + + +_CANCEL = {"jsonrpc": "2.0", "method": "$/cancel_request"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("request_id", "cancel"), + [ + (None, _CANCEL), + (None, {**_CANCEL, "params": None}), + (None, {**_CANCEL, "params": []}), + (None, {**_CANCEL, "params": {}}), + (1, _cancel_request(True)), + (0, _cancel_request(False)), + (1, _cancel_request(1.0)), + (1, _cancel_request([1])), + (1, _cancel_request("1")), + # Ids outside the schema's ``RequestId`` are not cancellable, not even by the integer they equal. + (True, _cancel_request(1)), + (1.0, _cancel_request(1)), + ], + ids=[ + "no_params", + "null_params", + "list_params", + "no_request_id", + "true_vs_1", + "false_vs_0", + "float_vs_1", + "list_vs_1", + "str_vs_1", + "1_vs_true_id", + "1_vs_float_id", + ], +) +async def test_cancel_request_without_a_matching_request_id_is_ignored(request_id: Any, cancel: dict[str, Any]) -> None: + response = await _cancel_while_pending(request_id, cancel) + + assert response == {"jsonrpc": "2.0", "id": request_id, "result": {"ok": True}} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_id", [None, 0, 1, "1"]) +async def test_cancel_request_with_a_matching_request_id_cancels(request_id: int | str | None) -> None: + response = await _cancel_while_pending(request_id, _cancel_request(request_id)) + + assert response["id"] == request_id + assert response["error"]["code"] == -32800 + + +@pytest.mark.asyncio +async def test_handler_may_answer_cancel_request_with_a_result(server) -> None: + started = asyncio.Event() + + class _PartialAgent(TestAgent): + async def prompt(self, session_id: str, prompt: list[Any], **kwargs: Any) -> Any: + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + return {"stopReason": "cancelled"} + + async with AgentSideConnection( + cast(Agent, _PartialAgent()), server.server_writer, server.server_reader, listening=True + ): + await _write(server.client_writer, _prompt(3)) + await asyncio.wait_for(started.wait(), timeout=1) + await _write(server.client_writer, _cancel_request(3)) + + assert await _read(server.client_reader) == {"jsonrpc": "2.0", "id": 3, "result": {"stopReason": "cancelled"}} + + +@pytest.mark.asyncio +async def test_internally_cancelled_handler_replies_request_cancelled(server) -> None: + class _InternallyCancelledAgent(TestAgent): + async def prompt(self, session_id: str, prompt: list[Any], **kwargs: Any) -> Any: + inner = asyncio.ensure_future(asyncio.Event().wait()) + inner.cancel() + return await inner + + async with AgentSideConnection( + cast(Agent, _InternallyCancelledAgent()), server.server_writer, server.server_reader, listening=True + ): + await _write(server.client_writer, _prompt(4)) + + response = await _read(server.client_reader) + assert response["id"] == 4 + assert response["error"]["code"] == -32800 + + +@pytest.mark.asyncio +async def test_close_does_not_reply_to_in_flight_requests(server) -> None: + agent = _BlockingAgent() + async with AgentSideConnection( + cast(Agent, agent), server.server_writer, server.server_reader, listening=True + ) as conn: + await _write(server.client_writer, _prompt(5)) + await asyncio.wait_for(agent.started.wait(), timeout=1) + + await conn.close() + + assert agent.cancelled == ["sess"] + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(server.client_reader.readline(), timeout=0.1) + + +@pytest.mark.asyncio +async def test_close_does_not_hang_when_handler_returns_on_cancellation(server) -> None: + started = asyncio.Event() + + class _PartialAgent(TestAgent): + async def prompt(self, session_id: str, prompt: list[Any], **kwargs: Any) -> Any: + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + return {"stopReason": "cancelled"} + + async with AgentSideConnection( + cast(Agent, _PartialAgent()), server.server_writer, server.server_reader, listening=True + ) as conn: + await _write(server.client_writer, _prompt(8)) + await asyncio.wait_for(started.wait(), timeout=1) + + closing = asyncio.ensure_future(conn.close()) + done, _ = await asyncio.wait({closing}, timeout=1) + assert closing in done + + +@pytest.mark.asyncio +async def test_cancelling_an_outgoing_request_sends_cancel_request(server) -> None: + async with AgentSideConnection( + cast(Agent, TestAgent()), server.server_writer, server.server_reader, listening=True + ) as conn: + request = asyncio.create_task( + conn.request_permission( + session_id="sess", + tool_call=ToolCallUpdate(tool_call_id="call-1"), + options=[PermissionOption(option_id="allow", name="Allow", kind="allow_once")], + ) + ) + + outgoing = await _read(server.client_reader) + assert outgoing["method"] == "session/request_permission" + request.cancel() + with pytest.raises(asyncio.CancelledError): + await request + + assert await _read(server.client_reader) == { + "jsonrpc": "2.0", + "method": "$/cancel_request", + "params": {"requestId": outgoing["id"]}, + } + # A late reply to the abandoned request is dropped and the connection stays usable. + await _write(server.client_writer, {"jsonrpc": "2.0", "id": outgoing["id"], "error": {"code": -32800}}) + await _write(server.client_writer, _prompt(6)) + assert await _read(server.client_reader) == {"jsonrpc": "2.0", "id": 6, "result": {"stopReason": "end_turn"}} + + +@pytest.mark.asyncio +async def test_completed_outgoing_request_does_not_send_cancel_request(server) -> None: + async with AgentSideConnection( + cast(Agent, TestAgent()), server.server_writer, server.server_reader, listening=True + ) as conn: + request = asyncio.create_task( + conn.request_permission( + session_id="sess", + tool_call=ToolCallUpdate(tool_call_id="call-1"), + options=[PermissionOption(option_id="allow", name="Allow", kind="allow_once")], + ) + ) + outgoing = await _read(server.client_reader) + await _write( + server.client_writer, + {"jsonrpc": "2.0", "id": outgoing["id"], "result": {"outcome": {"outcome": "cancelled"}}}, + ) + response = await asyncio.wait_for(request, timeout=1) + + assert response.outcome.outcome == "cancelled" + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(server.client_reader.readline(), timeout=0.1) + + +@pytest.mark.asyncio +async def test_cancellation_cascades_between_sdk_peers(server) -> None: + permission_started = asyncio.Event() + permission_cancelled = asyncio.Event() + + class _WaitingClient(TestClient): + async def request_permission(self, session_id: str, tool_call: Any, options: Any, **kwargs: Any) -> Any: + permission_started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + permission_cancelled.set() + raise + + async with ( + AgentSideConnection( + cast(Agent, TestAgent()), server.server_writer, server.server_reader, listening=True + ) as agent_conn, + ClientSideConnection(_WaitingClient(), server.client_writer, server.client_reader), + ): + request = asyncio.create_task( + agent_conn.request_permission( + session_id="sess", + tool_call=ToolCallUpdate(tool_call_id="call-1"), + options=[PermissionOption(option_id="allow", name="Allow", kind="allow_once")], + ) + ) + await asyncio.wait_for(permission_started.wait(), timeout=1) + + request.cancel() + with pytest.raises(asyncio.CancelledError): + await request + + await asyncio.wait_for(permission_cancelled.wait(), timeout=1) diff --git a/tests/test_v2_runtime.py b/tests/test_v2_runtime.py index 6300f5d..d93da0b 100644 --- a/tests/test_v2_runtime.py +++ b/tests/test_v2_runtime.py @@ -515,3 +515,70 @@ async def create_elicitation(self, message, mode, **kwargs): assert elicitations[2]["request_id"] is None with pytest.raises(ValueError, match="either session_id or request_id"): await agent_connection.create_elicitation("Input", "vendor/custom", session_id="s", request_id=7) + + +@pytest.mark.asyncio +async def test_cancelling_a_request_cancels_the_remote_handler() -> None: + started = asyncio.Event() + cancelled = asyncio.Event() + + class SlowAgent(ExtensionAgent): + async def handle_extension_request(self, method: str, params: Any) -> Any: + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cancelled.set() + raise + + client_transport, agent_transport = memory_transport_pair() + agent_connection = v2.AgentSideConnection(SlowAgent(), agent_transport) + client_connection = v2.ClientSideConnection(ExtensionClient(), client_transport) + + try: + await client_connection.initialize( + protocol_version=v2.PROTOCOL_VERSION, info=v2.schema.Implementation(name="test-client", version="1.0.0") + ) + request = asyncio.create_task(client_connection.send_extension_request("_vendor/slow")) + await asyncio.wait_for(started.wait(), timeout=1) + + request.cancel() + with pytest.raises(asyncio.CancelledError): + await request + + await asyncio.wait_for(cancelled.wait(), timeout=1) + finally: + await client_connection.close() + await agent_connection.close() + + +@pytest.mark.asyncio +async def test_cancel_request_before_the_handler_starts_still_replies() -> None: + handled: list[str] = [] + + class RecordingAgent(ExtensionAgent): + async def handle_extension_request(self, method: str, params: Any) -> Any: + handled.append(method) + return {} + + peer, agent_transport = memory_transport_pair() + agent_connection = v2.AgentSideConnection(RecordingAgent(), agent_transport) + + try: + await peer.send({ + "jsonrpc": "2.0", + "id": 0, + "method": "initialize", + "params": initialize_request().model_dump(mode="json", by_alias=True, exclude_none=True), + }) + assert "result" in await asyncio.wait_for(peer.receive(), timeout=1) + + # Both frames are queued before the connection runs, so the cancel precedes the handler. + await peer.send({"jsonrpc": "2.0", "id": 1, "method": "_vendor/slow", "params": {}}) + await peer.send({"jsonrpc": "2.0", "method": "$/cancel_request", "params": {"requestId": 1}}) + response = await asyncio.wait_for(peer.receive(), timeout=1) + + assert (response["id"], response["error"]["code"]) == (1, -32800) + assert handled == [] + finally: + await agent_connection.close()