1010import pytest
1111
1212from acp import Agent
13+ from acp .connection import Connection
1314from acp .core import AgentSideConnection , ClientSideConnection
1415from acp .schema import PermissionOption , ToolCallUpdate
1516from tests .conftest import TestAgent , TestClient
@@ -24,7 +25,7 @@ async def _read(reader: asyncio.StreamReader) -> dict[str, Any]:
2425 return json .loads (await asyncio .wait_for (reader .readline (), timeout = 1 ))
2526
2627
27- def _prompt (request_id : int | str ) -> dict [str , Any ]:
28+ def _prompt (request_id : int | str | None ) -> dict [str , Any ]:
2829 return {
2930 "jsonrpc" : "2.0" ,
3031 "id" : request_id ,
@@ -33,7 +34,7 @@ def _prompt(request_id: int | str) -> dict[str, Any]:
3334 }
3435
3536
36- def _cancel_request (request_id : int | str ) -> dict [str , Any ]:
37+ def _cancel_request (request_id : int | str | None ) -> dict [str , Any ]:
3738 return {"jsonrpc" : "2.0" , "method" : "$/cancel_request" , "params" : {"requestId" : request_id }}
3839
3940
@@ -77,6 +78,122 @@ async def test_cancel_request_cancels_handler_and_replies_request_cancelled(
7778 assert "$/cancel_request" not in caplog .text
7879
7980
81+ @pytest .mark .asyncio
82+ @pytest .mark .parametrize ("request_id" , [0 , "req-0" , None ])
83+ async def test_cancel_request_before_the_handler_starts_still_replies (
84+ server , caplog : pytest .LogCaptureFixture , request_id : int | str | None
85+ ) -> None :
86+ agent = _BlockingAgent ()
87+ async with AgentSideConnection (cast (Agent , agent ), server .server_writer , server .server_reader , listening = True ):
88+ with caplog .at_level (logging .ERROR ):
89+ # One write, so the receive loop reads both frames before the handler task first runs.
90+ server .client_writer .write (
91+ (json .dumps (_prompt (request_id )) + "\n " + json .dumps (_cancel_request (request_id )) + "\n " ).encode ()
92+ )
93+ await server .client_writer .drain ()
94+ response = await _read (server .client_reader )
95+
96+ assert response ["id" ] == request_id
97+ assert response ["error" ]["code" ] == - 32800
98+ assert not agent .started .is_set ()
99+ assert caplog .text == ""
100+ with pytest .raises (asyncio .TimeoutError ):
101+ await asyncio .wait_for (server .client_reader .readline (), timeout = 0.1 )
102+
103+
104+ class _GatedTransport :
105+ """Message transport whose sends block until ``release`` is set."""
106+
107+ def __init__ (self ) -> None :
108+ self .incoming : asyncio .Queue [dict [str , Any ]] = asyncio .Queue ()
109+ self .sent : list [dict [str , Any ]] = []
110+ self .sending = asyncio .Event ()
111+ self .release = asyncio .Event ()
112+ self .settled = asyncio .Event ()
113+ self .send_cancelled = False
114+ self ._receiving = False
115+
116+ async def send (self , message : dict [str , Any ]) -> None :
117+ self .sending .set ()
118+ try :
119+ await self .release .wait ()
120+ except asyncio .CancelledError :
121+ self .send_cancelled = True
122+ raise
123+ else :
124+ self .sent .append (message )
125+ finally :
126+ self .settled .set ()
127+
128+ async def receive (self ) -> dict [str , Any ] | None :
129+ self ._receiving = True
130+ try :
131+ return await self .incoming .get ()
132+ finally :
133+ self ._receiving = False
134+
135+ async def close (self ) -> None :
136+ pass
137+
138+ async def deliver (self , message : dict [str , Any ]) -> None :
139+ """Queue ``message`` and wait until the connection has processed it."""
140+ await self .incoming .put (message )
141+ for _ in range (100 ):
142+ if self ._receiving and self .incoming .empty ():
143+ return
144+ await asyncio .sleep (0 )
145+ raise AssertionError ("the connection did not process the message" )
146+
147+
148+ @pytest .mark .asyncio
149+ @pytest .mark .parametrize ("handler_cancelled" , [False , True ], ids = ["result" , "request_cancelled" ])
150+ async def test_cancel_request_during_response_send_keeps_the_response (handler_cancelled : bool ) -> None :
151+ transport = _GatedTransport ()
152+ started = asyncio .Event ()
153+
154+ async def handler (method : str , params : Any , is_notification : bool ) -> Any :
155+ started .set ()
156+ if handler_cancelled :
157+ await asyncio .Event ().wait ()
158+ return {"ok" : True }
159+
160+ async with Connection (handler , transport ):
161+ await transport .deliver (_prompt (0 ))
162+ await asyncio .wait_for (started .wait (), timeout = 1 )
163+ if handler_cancelled :
164+ await transport .deliver (_cancel_request (0 ))
165+ await asyncio .wait_for (transport .sending .wait (), timeout = 1 )
166+
167+ # A late (or repeated) cancellation lands while the response is still being sent.
168+ await transport .deliver (_cancel_request (0 ))
169+ assert transport .sent == []
170+ transport .release .set ()
171+ await asyncio .wait_for (transport .settled .wait (), timeout = 1 )
172+ assert not transport .send_cancelled , "the cancellation aborted the response send"
173+
174+ if handler_cancelled :
175+ assert [(m ["id" ], m ["error" ]["code" ]) for m in transport .sent ] == [(0 , - 32800 )]
176+ else :
177+ assert transport .sent == [{"jsonrpc" : "2.0" , "id" : 0 , "result" : {"ok" : True }}]
178+
179+
180+ @pytest .mark .asyncio
181+ async def test_close_cancels_a_blocked_response_send () -> None :
182+ transport = _GatedTransport ()
183+
184+ async def handler (method : str , params : Any , is_notification : bool ) -> Any :
185+ return {"ok" : True }
186+
187+ conn = Connection (handler , transport )
188+ await transport .deliver (_prompt (0 ))
189+ await asyncio .wait_for (transport .sending .wait (), timeout = 1 )
190+
191+ await asyncio .wait_for (conn .close (), timeout = 1 )
192+
193+ assert transport .send_cancelled
194+ assert transport .sent == []
195+
196+
80197@pytest .mark .asyncio
81198async def test_cancel_request_only_affects_the_targeted_request (server , caplog : pytest .LogCaptureFixture ) -> None :
82199 agent = _BlockingAgent ()
0 commit comments