1414
1515from ._transport import NdjsonTransport , Transport
1616from .exceptions import RequestError
17- from .task import (
18- DefaultMessageDispatcher ,
19- InMemoryMessageQueue ,
20- InMemoryMessageStateStore ,
21- MessageDispatcher ,
22- MessageQueue ,
23- MessageSender ,
24- MessageStateStore ,
25- NotificationRunner ,
26- RequestRunner ,
27- RpcTask ,
28- RpcTaskKind ,
29- SenderFactory ,
30- TaskSupervisor ,
31- )
17+ from .task import MessageSender , TaskSupervisor
3218from .telemetry import span_context
3319
3420JsonValue = Any
3824__all__ = ["Connection" , "JsonValue" , "MethodHandler" , "StreamDirection" , "StreamEvent" ]
3925
4026
41- DispatcherFactory = Callable [
42- [MessageQueue , TaskSupervisor , MessageStateStore , RequestRunner , NotificationRunner ],
43- MessageDispatcher ,
44- ]
45-
46-
4727class StreamDirection (str , Enum ):
4828 INCOMING = "incoming"
4929 OUTGOING = "outgoing"
@@ -67,20 +47,15 @@ def __init__(
6747 writer : asyncio .StreamWriter | Transport ,
6848 reader : asyncio .StreamReader | None = None ,
6949 * ,
70- queue : MessageQueue | None = None ,
71- state_store : MessageStateStore | None = None ,
72- dispatcher_factory : DispatcherFactory | None = None ,
73- sender_factory : SenderFactory | None = None ,
7450 observers : list [StreamObserver ] | None = None ,
7551 listening : bool = True ,
7652 receive_timeout : float | None = None ,
7753 ) -> None :
7854 self ._handler = handler
7955 self ._next_request_id = 0
80- self ._state = state_store or InMemoryMessageStateStore ()
56+ self ._pending : dict [ int , asyncio . Future [ Any ]] = {}
8157 self ._tasks = TaskSupervisor (source = "acp.Connection" )
8258 self ._tasks .add_error_handler (self ._on_task_error )
83- self ._queue = queue or InMemoryMessageQueue ()
8459 self ._closed = False
8560 self ._disconnected = False
8661 # Two construction forms:
@@ -92,7 +67,7 @@ def __init__(
9267 if reader is None :
9368 self ._transport : Transport = cast ("Transport" , writer )
9469 else :
95- sender = ( sender_factory or self . _default_sender_factory ) (cast ("asyncio.StreamWriter" , writer ), self ._tasks )
70+ sender = MessageSender (cast ("asyncio.StreamWriter" , writer ), self ._tasks )
9671 self ._transport = NdjsonTransport (reader , sender , receive_timeout = receive_timeout )
9772 self ._observers : list [StreamObserver ] = list (observers or [])
9873 if listening :
@@ -103,25 +78,17 @@ def __init__(
10378 )
10479 else :
10580 self ._recv_task = None
106- dispatcher_factory = dispatcher_factory or self ._default_dispatcher_factory
107- self ._dispatcher = dispatcher_factory (
108- self ._queue ,
109- self ._tasks ,
110- self ._state ,
111- self ._run_request ,
112- self ._run_notification ,
113- )
114- self ._dispatcher .start ()
11581
11682 async def close (self ) -> None :
11783 """Stop the receive loop and cancel any in-flight handler tasks."""
11884 if self ._closed :
11985 return
12086 self ._closed = True
121- await self ._dispatcher .stop ()
122- await self ._transport .close ()
123- await self ._tasks .shutdown ()
124- self ._state .reject_all_outgoing (ConnectionError ("Connection closed" ))
87+ self ._reject_all_outgoing (ConnectionError ("Connection closed" ))
88+ try :
89+ await self ._transport .close ()
90+ finally :
91+ await self ._tasks .shutdown ()
12592
12693 async def main_loop (self ) -> None :
12794 try :
@@ -145,18 +112,22 @@ async def send_request(self, method: str, params: JsonValue | None = None) -> An
145112 self ._raise_if_unavailable ()
146113 request_id = self ._next_request_id
147114 self ._next_request_id += 1
148- future = self ._state .register_outgoing (request_id , method )
115+ future : asyncio .Future [Any ] = asyncio .get_running_loop ().create_future ()
116+ self ._pending [request_id ] = future
149117 payload = {"jsonrpc" : "2.0" , "id" : request_id , "method" : method , "params" : params }
150118 try :
151119 await self ._transport .send (payload )
152- except Exception as exc :
153- # A synchronous send failure (e.g. HTTP POST rejected before any
154- # JSON-RPC response exists) must reject the correlated future so the
155- # caller gets a real, attributable error.
156- self ._state .reject_outgoing (request_id , exc )
120+ except BaseException :
121+ self ._pending .pop (request_id , None )
122+ future .cancel ()
157123 raise
158124 self ._notify_observers (StreamDirection .OUTGOING , payload )
159- return await future
125+ try :
126+ return await future
127+ except asyncio .CancelledError :
128+ self ._pending .pop (request_id , None )
129+ future .cancel ()
130+ raise
160131
161132 async def send_notification (self , method : str , params : JsonValue | None = None ) -> None :
162133 self ._raise_if_unavailable ()
@@ -171,24 +142,26 @@ async def _receive_loop(self) -> None:
171142 if message is None :
172143 break
173144 self ._notify_observers (StreamDirection .INCOMING , message )
174- await self ._process_message (message )
145+ self ._process_message (message )
175146 except asyncio .CancelledError :
176147 return
177148 except asyncio .TimeoutError :
178149 raise RequestError .internal_error ({"details" : "Agent timeout" }) from None
179150 self ._disconnect ()
180151
181- async def _process_message (self , message : dict [str , Any ]) -> None :
152+ def _process_message (self , message : dict [str , Any ]) -> None :
182153 method = message .get ("method" )
183154 has_id = "id" in message
184- if method is not None and has_id :
185- await self ._queue .publish (RpcTask (RpcTaskKind .REQUEST , message ))
186- return
187- if method is not None and not has_id :
188- await self ._queue .publish (RpcTask (RpcTaskKind .NOTIFICATION , message ))
155+ if method is not None : # this is a request or notification
156+ # {"jsonrpc": "2.0", "id": 1, "method": "foo", "params": {...}} # request
157+ # {"jsonrpc": "2.0", "method": "foo", "params: {...}} # notification
158+ self ._tasks .create (
159+ self ._run_request (message ) if has_id else self ._run_notification (message ),
160+ name = "acp.Connection.request" if has_id else "acp.Connection.notification" ,
161+ )
189162 return
190- if has_id :
191- await self ._handle_response (message )
163+ if has_id : # this is a response, {"id", "result" | "error"}
164+ self ._handle_response (message )
192165
193166 def _notify_observers (self , direction : StreamDirection , message : dict [str , Any ]) -> None :
194167 if not self ._observers :
@@ -211,7 +184,12 @@ def _notify_observers(self, direction: StreamDirection, message: dict[str, Any])
211184 def _on_observer_error (self , task : asyncio .Task [Any ], exc : BaseException ) -> None :
212185 logging .exception ("Stream observer coroutine failed" , exc_info = exc )
213186
214- async def _run_request (self , message : dict [str , Any ]) -> Any :
187+ async def _run_request (self , message : dict [str , Any ]) -> None :
188+ payload = await self ._execute_request (message )
189+ await self ._transport .send (payload )
190+ self ._notify_observers (StreamDirection .OUTGOING , payload )
191+
192+ async def _execute_request (self , message : dict [str , Any ]) -> dict [str , Any ]:
215193 payload : dict [str , Any ] = {"jsonrpc" : "2.0" , "id" : message ["id" ]}
216194 method = message ["method" ]
217195 with span_context (
@@ -228,20 +206,10 @@ async def _run_request(self, message: dict[str, Any]) -> Any:
228206 exclude_unset = True ,
229207 )
230208 payload ["result" ] = result if result is not None else None
231- await self ._transport .send (payload )
232- self ._notify_observers (StreamDirection .OUTGOING , payload )
233- return payload .get ("result" )
234209 except RequestError as exc :
235210 payload ["error" ] = exc .to_error_obj ()
236- await self ._transport .send (payload )
237- self ._notify_observers (StreamDirection .OUTGOING , payload )
238- raise
239211 except ValidationError as exc :
240- err = RequestError .invalid_params ({"errors" : exc .errors ()})
241- payload ["error" ] = err .to_error_obj ()
242- await self ._transport .send (payload )
243- self ._notify_observers (StreamDirection .OUTGOING , payload )
244- raise err from None
212+ payload ["error" ] = RequestError .invalid_params ({"errors" : exc .errors ()}).to_error_obj ()
245213 except Exception as exc :
246214 logging .exception (
247215 "Unhandled error while handling request method=%s" ,
@@ -252,11 +220,8 @@ async def _run_request(self, message: dict[str, Any]) -> Any:
252220 data = json .loads (str (exc ))
253221 except Exception :
254222 data = {"details" : str (exc )}
255- err = RequestError .internal_error (data )
256- payload ["error" ] = err .to_error_obj ()
257- await self ._transport .send (payload )
258- self ._notify_observers (StreamDirection .OUTGOING , payload )
259- raise err from None
223+ payload ["error" ] = RequestError .internal_error (data ).to_error_obj ()
224+ return payload
260225
261226 async def _run_notification (self , message : dict [str , Any ]) -> None :
262227 method = message ["method" ]
@@ -270,24 +235,21 @@ async def _run_notification(self, message: dict[str, Any]) -> None:
270235 exc_info = exc ,
271236 )
272237
273- async def _handle_response (self , message : dict [str , Any ]) -> None :
238+ def _handle_response (self , message : dict [str , Any ]) -> None :
274239 request_id = message ["id" ]
275- result = message .get ("result" )
240+ future = self ._pending .pop (request_id , None )
241+ if future is None or future .done ():
242+ return
276243 if "result" in message :
277- self . _state . resolve_outgoing ( request_id , result )
244+ future . set_result ( message . get ( " result" ) )
278245 return
279246 if "error" in message :
280247 error_obj = message .get ("error" ) or {}
281- self ._state .reject_outgoing (
282- request_id ,
283- RequestError (
284- error_obj .get ("code" , - 32603 ),
285- error_obj .get ("message" , "Error" ),
286- error_obj .get ("data" ),
287- ),
248+ future .set_exception (
249+ RequestError (error_obj .get ("code" , - 32603 ), error_obj .get ("message" , "Error" ), error_obj .get ("data" ))
288250 )
289251 return
290- self . _state . resolve_outgoing ( request_id , None )
252+ future . set_result ( None )
291253
292254 def _on_receive_error (self , task : asyncio .Task [Any ], exc : BaseException ) -> None :
293255 logging .exception ("Receive loop failed" , exc_info = exc )
@@ -296,30 +258,18 @@ def _on_receive_error(self, task: asyncio.Task[Any], exc: BaseException) -> None
296258 def _on_task_error (self , task : asyncio .Task [Any ], exc : BaseException ) -> None :
297259 logging .exception ("Background task failed" , exc_info = exc )
298260
299- def _default_dispatcher_factory (
300- self ,
301- queue : MessageQueue ,
302- supervisor : TaskSupervisor ,
303- state : MessageStateStore ,
304- request_runner : RequestRunner ,
305- notification_runner : NotificationRunner ,
306- ) -> MessageDispatcher :
307- return DefaultMessageDispatcher (
308- queue = queue ,
309- supervisor = supervisor ,
310- store = state ,
311- request_runner = request_runner ,
312- notification_runner = notification_runner ,
313- )
314-
315- def _default_sender_factory (self , writer : asyncio .StreamWriter , supervisor : TaskSupervisor ) -> MessageSender :
316- return MessageSender (writer , supervisor )
317-
318261 def _disconnect (self ) -> None :
319262 if self ._disconnected :
320263 return
321264 self ._disconnected = True
322- self ._state .reject_all_outgoing (ConnectionError ("Connection closed" ))
265+ self ._reject_all_outgoing (ConnectionError ("Connection closed" ))
266+
267+ def _reject_all_outgoing (self , error : BaseException ) -> None :
268+ pending = list (self ._pending .values ())
269+ self ._pending .clear ()
270+ for future in pending :
271+ if not future .done ():
272+ future .set_exception (error )
323273
324274 def _raise_if_unavailable (self ) -> None :
325275 if self ._disconnected or self ._closed :
0 commit comments