From e74cbc6904bcbec1ebf309ebe39770791c3161c6 Mon Sep 17 00:00:00 2001 From: monody0007 <52037177+monody0007@users.noreply.github.com> Date: Mon, 28 Sep 2026 10:00:56 -0700 Subject: [PATCH 1/3] fix(connection): honor $/cancel_request and reply -32800 to cancelled requests ACP v1 defines the `$/cancel_request` notification (part of the stable v1 schema vendored in schema/, 1.23.0; see https://agentclientprotocol.com/protocol/v1/cancellation), and the SDK already ships `PROTOCOL_METHODS["cancel_request"]`, but `Connection` never handled it: - An incoming `$/cancel_request` fell through to the router, which logged an ERROR traceback (method not found) while the targeted handler kept running and the peer never received the required response. - A handler that ended with `CancelledError` (cancelled from inside the agent) produced no response at all, leaving the peer waiting forever. - Cancelling the task awaiting `send_request` only dropped the local future; the peer kept working on the request. - `close()` hung forever when a handler answered its own cancellation with a result, because the response was queued on an already-closed sender. `Connection` now tracks in-flight incoming requests by id, cancels the handler task on `$/cancel_request`, answers cancelled requests with `-32800 Request cancelled` (a handler may still return a partial result), sends `$/cancel_request` when a still-pending outgoing request is cancelled locally, and skips replies once the connection is closed. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/experimental-v2.md | 6 +- src/acp/connection.py | 76 ++++++-- src/acp/exceptions.py | 4 + tests/test_request_cancellation.py | 270 +++++++++++++++++++++++++++++ tests/test_v2_runtime.py | 35 ++++ 5 files changed, 379 insertions(+), 12 deletions(-) create mode 100644 tests/test_request_cancellation.py diff --git a/docs/experimental-v2.md b/docs/experimental-v2.md index 7c1149a..62dfa89 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 +transport. Protocol-level request cancellation (`$/cancel_request`) is handled +by the shared connection layer, as in v1: cancelling the task awaiting a request +notifies the peer, and an incoming cancellation cancels the handler task and +answers with `-32800`. `session/cancel` remains available for cancelling active session work. diff --git a/src/acp/connection.py b/src/acp/connection.py index 84f8f05..83c38fb 100644 --- a/src/acp/connection.py +++ b/src/acp/connection.py @@ -5,6 +5,7 @@ import inspect import json import logging +import sys from collections.abc import Awaitable, Callable from dataclasses import dataclass from enum import Enum @@ -14,6 +15,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 +39,8 @@ class StreamEvent: StreamObserver = Callable[[StreamEvent], Awaitable[None] | None] +_CANCEL_REQUEST_METHOD = PROTOCOL_METHODS["cancel_request"] + class Connection: """Minimal JSON-RPC 2.0 connection over newline-delimited JSON frames.""" @@ -54,6 +58,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 +122,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 +155,18 @@ 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( + task = 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 has_id: + self._track_incoming(message["id"], task) return if has_id: # this is a response, {"id", "result" | "error"} self._handle_response(message) @@ -184,8 +192,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) + def _track_incoming(self, request_id: Any, task: asyncio.Task[Any]) -> None: + try: + self._incoming[request_id] = task + except TypeError: # unhashable id, nothing can refer to it + return + + 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: + request_id = params.get("requestId") if isinstance(params, dict) else None + try: + task = self._incoming.get(request_id) + except TypeError: + return + # Unknown or already finished requests are ignored, as the protocol allows. + 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]) -> None: - payload = await self._execute_request(message) + try: + payload = await self._execute_request(message) + except asyncio.CancelledError: + if self._closed: + raise + # Cancelled by ``$/cancel_request`` or from inside the handler; either way the + # protocol still requires a response for the original request. + task = asyncio.current_task() + if sys.version_info >= (3, 11) and task is not None: + task.uncancel() + payload = {"jsonrpc": "2.0", "id": message["id"], "error": RequestError.request_cancelled().to_error_obj()} + if self._closed: + # The transport is gone (e.g. the handler returned a result while being + # cancelled by ``close()``); sending would never complete. + return 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..523c1be --- /dev/null +++ b/tests/test_request_cancellation.py @@ -0,0 +1,270 @@ +"""``$/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.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: int | str) -> 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: int | str) -> 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 +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) + + +@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..ab13a1c 100644 --- a/tests/test_v2_runtime.py +++ b/tests/test_v2_runtime.py @@ -515,3 +515,38 @@ 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() From e1abcac1345400d78463ae0b1ddf1002f3d7d1d3 Mon Sep 17 00:00:00 2001 From: monody0007 <52037177+monody0007@users.noreply.github.com> Date: Mon, 28 Sep 2026 10:35:28 -0700 Subject: [PATCH 2/3] fix(connection): keep request responses out of the handler's cancellation Two races could still lose the response to a cancelled request: - A `$/cancel_request` read in the same buffer as its request cancelled the request task before it first ran, so the coroutine that would have sent `-32800` never executed. - A late or repeated `$/cancel_request` arriving while the response was being sent cancelled that send, so the peer got no response at all. Each incoming request now runs its handler in its own task, which is the only thing `$/cancel_request` targets. A separate supervised task waits for the handler with `asyncio.wait` and sends its result, or `-32800` if the handler was cancelled. `close()` still cancels both, so shutdown never waits on a blocked send, and the `Task.uncancel()` / closed-send special cases are no longer needed. Also document request cancellation in the quickstart and state that a cancelled handler answers with its result or `-32800`. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/experimental-v2.md | 6 +- docs/quickstart.md | 15 ++++ src/acp/connection.py | 42 +++++----- tests/test_request_cancellation.py | 121 ++++++++++++++++++++++++++++- tests/test_v2_runtime.py | 32 ++++++++ 5 files changed, 189 insertions(+), 27 deletions(-) diff --git a/docs/experimental-v2.md b/docs/experimental-v2.md index 62dfa89..e62ecb6 100644 --- a/docs/experimental-v2.md +++ b/docs/experimental-v2.md @@ -152,6 +152,6 @@ when an agent selects v1. Client-side fallback is application controlled and may require opening a new transport. Protocol-level request cancellation (`$/cancel_request`) is handled by the shared connection layer, as in v1: cancelling the task awaiting a request -notifies the peer, and an incoming cancellation cancels the handler task and -answers with `-32800`. `session/cancel` remains available for cancelling active -session work. +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 83c38fb..82ef48f 100644 --- a/src/acp/connection.py +++ b/src/acp/connection.py @@ -5,7 +5,6 @@ import inspect import json import logging -import sys from collections.abc import Awaitable, Callable from dataclasses import dataclass from enum import Enum @@ -161,12 +160,14 @@ def _process_message(self, message: dict[str, Any]) -> None: 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 - task = 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 has_id: - self._track_incoming(message["id"], task) + 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) @@ -226,22 +227,19 @@ def _send_cancel_request(self, request_id: int) -> None: 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]) -> None: - try: - payload = await self._execute_request(message) - except asyncio.CancelledError: - if self._closed: - raise - # Cancelled by ``$/cancel_request`` or from inside the handler; either way the - # protocol still requires a response for the original request. - task = asyncio.current_task() - if sys.version_info >= (3, 11) and task is not None: - task.uncancel() + 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()} - if self._closed: - # The transport is gone (e.g. the handler returned a result while being - # cancelled by ``close()``); sending would never complete. - return + else: + payload = handler.result() await self._transport.send(payload) self._notify_observers(StreamDirection.OUTGOING, payload) diff --git a/tests/test_request_cancellation.py b/tests/test_request_cancellation.py index 523c1be..462e72a 100644 --- a/tests/test_request_cancellation.py +++ b/tests/test_request_cancellation.py @@ -10,6 +10,7 @@ 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 @@ -24,7 +25,7 @@ async def _read(reader: asyncio.StreamReader) -> dict[str, Any]: return json.loads(await asyncio.wait_for(reader.readline(), timeout=1)) -def _prompt(request_id: int | str) -> dict[str, Any]: +def _prompt(request_id: int | str | None) -> dict[str, Any]: return { "jsonrpc": "2.0", "id": request_id, @@ -33,7 +34,7 @@ def _prompt(request_id: int | str) -> dict[str, Any]: } -def _cancel_request(request_id: int | str) -> dict[str, Any]: +def _cancel_request(request_id: int | str | None) -> dict[str, Any]: return {"jsonrpc": "2.0", "method": "$/cancel_request", "params": {"requestId": request_id}} @@ -77,6 +78,122 @@ async def test_cancel_request_cancels_handler_and_replies_request_cancelled( 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() diff --git a/tests/test_v2_runtime.py b/tests/test_v2_runtime.py index ab13a1c..d93da0b 100644 --- a/tests/test_v2_runtime.py +++ b/tests/test_v2_runtime.py @@ -550,3 +550,35 @@ async def handle_extension_request(self, method: str, params: Any) -> Any: 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() From e03a3216106fde9f9386d1bde8fc57ac4d407a59 Mon Sep 17 00:00:00 2001 From: monody0007 <52037177+monody0007@users.noreply.github.com> Date: Mon, 28 Sep 2026 18:26:04 -0700 Subject: [PATCH 3/3] fix(connection): only cancel requests named by a valid requestId `$/cancel_request` requires a `requestId`, and a request id is `null`, an integer or a string. The handler did not check either, so a malformed cancellation could abort an unrelated request: - A notification without `params`, with non-object `params` or without `requestId` was treated as `requestId: null` and cancelled a pending request whose id is `null`. - `requestId: true` / `false` cancelled pending requests `1` / `0`, and `1.0` cancelled `1`, because Python equates these dict keys. Likewise a request sent with a non-integer id such as `true` or `1.0` could be cancelled by a cancellation for `1`. Cancellations without a `requestId`, or with one that is not a valid request id, are now ignored, and only requests with a valid id are tracked for cancellation. An explicit `requestId: null` still cancels a `null` request. --- src/acp/connection.py | 22 ++++++--- tests/test_request_cancellation.py | 76 +++++++++++++++++++++++++++++- 2 files changed, 89 insertions(+), 9 deletions(-) diff --git a/src/acp/connection.py b/src/acp/connection.py index 82ef48f..a827210 100644 --- a/src/acp/connection.py +++ b/src/acp/connection.py @@ -41,6 +41,11 @@ class StreamEvent: _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.""" @@ -194,10 +199,11 @@ def _on_observer_error(self, task: asyncio.Task[Any], exc: BaseException) -> Non logging.exception("Stream observer coroutine failed", exc_info=exc) def _track_incoming(self, request_id: Any, task: asyncio.Task[Any]) -> None: - try: - self._incoming[request_id] = task - except TypeError: # unhashable id, nothing can refer to it + # 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: @@ -206,12 +212,14 @@ def _forget(done: asyncio.Task[Any]) -> None: task.add_done_callback(_forget) def _cancel_incoming(self, params: Any) -> None: - request_id = params.get("requestId") if isinstance(params, dict) else None - try: - task = self._incoming.get(request_id) - except TypeError: + # ``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() diff --git a/tests/test_request_cancellation.py b/tests/test_request_cancellation.py index 462e72a..2206e7e 100644 --- a/tests/test_request_cancellation.py +++ b/tests/test_request_cancellation.py @@ -25,7 +25,7 @@ async def _read(reader: asyncio.StreamReader) -> dict[str, Any]: return json.loads(await asyncio.wait_for(reader.readline(), timeout=1)) -def _prompt(request_id: int | str | None) -> dict[str, Any]: +def _prompt(request_id: Any) -> dict[str, Any]: return { "jsonrpc": "2.0", "id": request_id, @@ -34,7 +34,7 @@ def _prompt(request_id: int | str | None) -> dict[str, Any]: } -def _cancel_request(request_id: int | str | None) -> dict[str, Any]: +def _cancel_request(request_id: Any) -> dict[str, Any]: return {"jsonrpc": "2.0", "method": "$/cancel_request", "params": {"requestId": request_id}} @@ -218,6 +218,78 @@ async def test_cancel_request_only_affects_the_targeted_request(server, caplog: 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()