Skip to content

Commit e7e96a7

Browse files
authored
refactor(connection): simplify connection handling and remove unused components (#132)
refactor(task): clean up task module by removing unused classes and imports test(tests): update tests to reflect changes in connection and task handling Signed-off-by: Frost Ming <me@frostming.com>
1 parent 082c8f0 commit e7e96a7

10 files changed

Lines changed: 148 additions & 456 deletions

File tree

src/acp/connection.py

Lines changed: 55 additions & 105 deletions
Original file line numberDiff line numberDiff line change
@@ -14,21 +14,7 @@
1414

1515
from ._transport import NdjsonTransport, Transport
1616
from .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
3218
from .telemetry import span_context
3319

3420
JsonValue = Any
@@ -38,12 +24,6 @@
3824
__all__ = ["Connection", "JsonValue", "MethodHandler", "StreamDirection", "StreamEvent"]
3925

4026

41-
DispatcherFactory = Callable[
42-
[MessageQueue, TaskSupervisor, MessageStateStore, RequestRunner, NotificationRunner],
43-
MessageDispatcher,
44-
]
45-
46-
4727
class 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:

src/acp/task/__init__.py

Lines changed: 3 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -1,44 +1,4 @@
1-
from __future__ import annotations
1+
from .sender import MessageSender
2+
from .supervisor import TaskSupervisor
23

3-
from dataclasses import dataclass
4-
from enum import Enum
5-
from typing import Any
6-
7-
__all__ = ["RpcTask", "RpcTaskKind"]
8-
9-
10-
class RpcTaskKind(Enum):
11-
REQUEST = "request"
12-
NOTIFICATION = "notification"
13-
14-
15-
@dataclass(slots=True)
16-
class RpcTask:
17-
kind: RpcTaskKind
18-
message: dict[str, Any]
19-
20-
21-
from .dispatcher import ( # noqa: E402
22-
DefaultMessageDispatcher,
23-
MessageDispatcher,
24-
NotificationRunner,
25-
RequestRunner,
26-
)
27-
from .queue import InMemoryMessageQueue, MessageQueue # noqa: E402
28-
from .sender import MessageSender, SenderFactory # noqa: E402
29-
from .state import InMemoryMessageStateStore, MessageStateStore # noqa: E402
30-
from .supervisor import TaskSupervisor # noqa: E402
31-
32-
__all__ += [
33-
"DefaultMessageDispatcher",
34-
"InMemoryMessageQueue",
35-
"InMemoryMessageStateStore",
36-
"MessageDispatcher",
37-
"MessageQueue",
38-
"MessageSender",
39-
"MessageStateStore",
40-
"NotificationRunner",
41-
"RequestRunner",
42-
"SenderFactory",
43-
"TaskSupervisor",
44-
]
4+
__all__ = ["MessageSender", "TaskSupervisor"]

0 commit comments

Comments
 (0)