From 1b912e1bf329fc3a7c828c238db6b5a029d2d609 Mon Sep 17 00:00:00 2001 From: Frost Ming Date: Thu, 13 Aug 2026 08:56:04 +0800 Subject: [PATCH] refactor(connection): simplify connection handling and remove unused components 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 --- src/acp/connection.py | 160 ++++++++++------------------ src/acp/task/__init__.py | 46 +------- src/acp/task/dispatcher.py | 94 ---------------- src/acp/task/queue.py | 67 ------------ src/acp/task/sender.py | 6 +- src/acp/task/state.py | 84 --------------- tests/test_connection_recovery.py | 4 +- tests/test_core.py | 38 +++---- tests/test_request_error_logging.py | 74 +++++++------ tests/test_rpc.py | 31 ++++++ 10 files changed, 148 insertions(+), 456 deletions(-) delete mode 100644 src/acp/task/dispatcher.py delete mode 100644 src/acp/task/queue.py delete mode 100644 src/acp/task/state.py diff --git a/src/acp/connection.py b/src/acp/connection.py index cfd7b5c..84f8f05 100644 --- a/src/acp/connection.py +++ b/src/acp/connection.py @@ -14,21 +14,7 @@ from ._transport import NdjsonTransport, Transport from .exceptions import RequestError -from .task import ( - DefaultMessageDispatcher, - InMemoryMessageQueue, - InMemoryMessageStateStore, - MessageDispatcher, - MessageQueue, - MessageSender, - MessageStateStore, - NotificationRunner, - RequestRunner, - RpcTask, - RpcTaskKind, - SenderFactory, - TaskSupervisor, -) +from .task import MessageSender, TaskSupervisor from .telemetry import span_context JsonValue = Any @@ -38,12 +24,6 @@ __all__ = ["Connection", "JsonValue", "MethodHandler", "StreamDirection", "StreamEvent"] -DispatcherFactory = Callable[ - [MessageQueue, TaskSupervisor, MessageStateStore, RequestRunner, NotificationRunner], - MessageDispatcher, -] - - class StreamDirection(str, Enum): INCOMING = "incoming" OUTGOING = "outgoing" @@ -67,20 +47,15 @@ def __init__( writer: asyncio.StreamWriter | Transport, reader: asyncio.StreamReader | None = None, *, - queue: MessageQueue | None = None, - state_store: MessageStateStore | None = None, - dispatcher_factory: DispatcherFactory | None = None, - sender_factory: SenderFactory | None = None, observers: list[StreamObserver] | None = None, listening: bool = True, receive_timeout: float | None = None, ) -> None: self._handler = handler self._next_request_id = 0 - self._state = state_store or InMemoryMessageStateStore() + self._pending: dict[int, asyncio.Future[Any]] = {} self._tasks = TaskSupervisor(source="acp.Connection") self._tasks.add_error_handler(self._on_task_error) - self._queue = queue or InMemoryMessageQueue() self._closed = False self._disconnected = False # Two construction forms: @@ -92,7 +67,7 @@ def __init__( if reader is None: self._transport: Transport = cast("Transport", writer) else: - sender = (sender_factory or self._default_sender_factory)(cast("asyncio.StreamWriter", writer), self._tasks) + sender = MessageSender(cast("asyncio.StreamWriter", writer), self._tasks) self._transport = NdjsonTransport(reader, sender, receive_timeout=receive_timeout) self._observers: list[StreamObserver] = list(observers or []) if listening: @@ -103,25 +78,17 @@ def __init__( ) else: self._recv_task = None - dispatcher_factory = dispatcher_factory or self._default_dispatcher_factory - self._dispatcher = dispatcher_factory( - self._queue, - self._tasks, - self._state, - self._run_request, - self._run_notification, - ) - self._dispatcher.start() async def close(self) -> None: """Stop the receive loop and cancel any in-flight handler tasks.""" if self._closed: return self._closed = True - await self._dispatcher.stop() - await self._transport.close() - await self._tasks.shutdown() - self._state.reject_all_outgoing(ConnectionError("Connection closed")) + self._reject_all_outgoing(ConnectionError("Connection closed")) + try: + await self._transport.close() + finally: + await self._tasks.shutdown() async def main_loop(self) -> None: try: @@ -145,18 +112,22 @@ async def send_request(self, method: str, params: JsonValue | None = None) -> An self._raise_if_unavailable() request_id = self._next_request_id self._next_request_id += 1 - future = self._state.register_outgoing(request_id, method) + future: asyncio.Future[Any] = asyncio.get_running_loop().create_future() + self._pending[request_id] = future payload = {"jsonrpc": "2.0", "id": request_id, "method": method, "params": params} try: await self._transport.send(payload) - except Exception as exc: - # A synchronous send failure (e.g. HTTP POST rejected before any - # JSON-RPC response exists) must reject the correlated future so the - # caller gets a real, attributable error. - self._state.reject_outgoing(request_id, exc) + except BaseException: + self._pending.pop(request_id, None) + future.cancel() raise self._notify_observers(StreamDirection.OUTGOING, payload) - return await future + try: + return await future + except asyncio.CancelledError: + self._pending.pop(request_id, None) + future.cancel() + raise async def send_notification(self, method: str, params: JsonValue | None = None) -> None: self._raise_if_unavailable() @@ -171,24 +142,26 @@ async def _receive_loop(self) -> None: if message is None: break self._notify_observers(StreamDirection.INCOMING, message) - await self._process_message(message) + self._process_message(message) except asyncio.CancelledError: return except asyncio.TimeoutError: raise RequestError.internal_error({"details": "Agent timeout"}) from None self._disconnect() - async def _process_message(self, message: dict[str, Any]) -> None: + def _process_message(self, message: dict[str, Any]) -> None: method = message.get("method") has_id = "id" in message - if method is not None and has_id: - await self._queue.publish(RpcTask(RpcTaskKind.REQUEST, message)) - return - if method is not None and not has_id: - await self._queue.publish(RpcTask(RpcTaskKind.NOTIFICATION, message)) + 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", + ) return - if has_id: - await self._handle_response(message) + if has_id: # this is a response, {"id", "result" | "error"} + self._handle_response(message) def _notify_observers(self, direction: StreamDirection, message: dict[str, Any]) -> None: if not self._observers: @@ -211,7 +184,12 @@ 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]) -> Any: + async def _run_request(self, message: dict[str, Any]) -> None: + payload = await self._execute_request(message) + await self._transport.send(payload) + self._notify_observers(StreamDirection.OUTGOING, payload) + + async def _execute_request(self, message: dict[str, Any]) -> dict[str, Any]: payload: dict[str, Any] = {"jsonrpc": "2.0", "id": message["id"]} method = message["method"] with span_context( @@ -228,20 +206,10 @@ async def _run_request(self, message: dict[str, Any]) -> Any: exclude_unset=True, ) payload["result"] = result if result is not None else None - await self._transport.send(payload) - self._notify_observers(StreamDirection.OUTGOING, payload) - return payload.get("result") except RequestError as exc: payload["error"] = exc.to_error_obj() - await self._transport.send(payload) - self._notify_observers(StreamDirection.OUTGOING, payload) - raise except ValidationError as exc: - err = RequestError.invalid_params({"errors": exc.errors()}) - payload["error"] = err.to_error_obj() - await self._transport.send(payload) - self._notify_observers(StreamDirection.OUTGOING, payload) - raise err from None + payload["error"] = RequestError.invalid_params({"errors": exc.errors()}).to_error_obj() except Exception as exc: logging.exception( "Unhandled error while handling request method=%s", @@ -252,11 +220,8 @@ async def _run_request(self, message: dict[str, Any]) -> Any: data = json.loads(str(exc)) except Exception: data = {"details": str(exc)} - err = RequestError.internal_error(data) - payload["error"] = err.to_error_obj() - await self._transport.send(payload) - self._notify_observers(StreamDirection.OUTGOING, payload) - raise err from None + payload["error"] = RequestError.internal_error(data).to_error_obj() + return payload async def _run_notification(self, message: dict[str, Any]) -> None: method = message["method"] @@ -270,24 +235,21 @@ async def _run_notification(self, message: dict[str, Any]) -> None: exc_info=exc, ) - async def _handle_response(self, message: dict[str, Any]) -> None: + def _handle_response(self, message: dict[str, Any]) -> None: request_id = message["id"] - result = message.get("result") + future = self._pending.pop(request_id, None) + if future is None or future.done(): + return if "result" in message: - self._state.resolve_outgoing(request_id, result) + future.set_result(message.get("result")) return if "error" in message: error_obj = message.get("error") or {} - self._state.reject_outgoing( - request_id, - RequestError( - error_obj.get("code", -32603), - error_obj.get("message", "Error"), - error_obj.get("data"), - ), + future.set_exception( + RequestError(error_obj.get("code", -32603), error_obj.get("message", "Error"), error_obj.get("data")) ) return - self._state.resolve_outgoing(request_id, None) + future.set_result(None) def _on_receive_error(self, task: asyncio.Task[Any], exc: BaseException) -> None: logging.exception("Receive loop failed", exc_info=exc) @@ -296,30 +258,18 @@ def _on_receive_error(self, task: asyncio.Task[Any], exc: BaseException) -> None def _on_task_error(self, task: asyncio.Task[Any], exc: BaseException) -> None: logging.exception("Background task failed", exc_info=exc) - def _default_dispatcher_factory( - self, - queue: MessageQueue, - supervisor: TaskSupervisor, - state: MessageStateStore, - request_runner: RequestRunner, - notification_runner: NotificationRunner, - ) -> MessageDispatcher: - return DefaultMessageDispatcher( - queue=queue, - supervisor=supervisor, - store=state, - request_runner=request_runner, - notification_runner=notification_runner, - ) - - def _default_sender_factory(self, writer: asyncio.StreamWriter, supervisor: TaskSupervisor) -> MessageSender: - return MessageSender(writer, supervisor) - def _disconnect(self) -> None: if self._disconnected: return self._disconnected = True - self._state.reject_all_outgoing(ConnectionError("Connection closed")) + self._reject_all_outgoing(ConnectionError("Connection closed")) + + def _reject_all_outgoing(self, error: BaseException) -> None: + pending = list(self._pending.values()) + self._pending.clear() + for future in pending: + if not future.done(): + future.set_exception(error) def _raise_if_unavailable(self) -> None: if self._disconnected or self._closed: diff --git a/src/acp/task/__init__.py b/src/acp/task/__init__.py index 2896fbf..074a67f 100644 --- a/src/acp/task/__init__.py +++ b/src/acp/task/__init__.py @@ -1,44 +1,4 @@ -from __future__ import annotations +from .sender import MessageSender +from .supervisor import TaskSupervisor -from dataclasses import dataclass -from enum import Enum -from typing import Any - -__all__ = ["RpcTask", "RpcTaskKind"] - - -class RpcTaskKind(Enum): - REQUEST = "request" - NOTIFICATION = "notification" - - -@dataclass(slots=True) -class RpcTask: - kind: RpcTaskKind - message: dict[str, Any] - - -from .dispatcher import ( # noqa: E402 - DefaultMessageDispatcher, - MessageDispatcher, - NotificationRunner, - RequestRunner, -) -from .queue import InMemoryMessageQueue, MessageQueue # noqa: E402 -from .sender import MessageSender, SenderFactory # noqa: E402 -from .state import InMemoryMessageStateStore, MessageStateStore # noqa: E402 -from .supervisor import TaskSupervisor # noqa: E402 - -__all__ += [ - "DefaultMessageDispatcher", - "InMemoryMessageQueue", - "InMemoryMessageStateStore", - "MessageDispatcher", - "MessageQueue", - "MessageSender", - "MessageStateStore", - "NotificationRunner", - "RequestRunner", - "SenderFactory", - "TaskSupervisor", -] +__all__ = ["MessageSender", "TaskSupervisor"] diff --git a/src/acp/task/dispatcher.py b/src/acp/task/dispatcher.py deleted file mode 100644 index e8c5e76..0000000 --- a/src/acp/task/dispatcher.py +++ /dev/null @@ -1,94 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections.abc import Awaitable, Callable -from contextlib import suppress -from typing import Any, Protocol - -from . import RpcTaskKind -from .queue import MessageQueue -from .state import MessageStateStore -from .supervisor import TaskSupervisor - -__all__ = [ - "DefaultMessageDispatcher", - "MessageDispatcher", - "NotificationRunner", - "RequestRunner", -] - - -RequestRunner = Callable[[dict[str, Any]], Awaitable[Any]] -NotificationRunner = Callable[[dict[str, Any]], Awaitable[None]] - - -class MessageDispatcher(Protocol): - def start(self) -> None: ... - - async def stop(self) -> None: ... - - -class DefaultMessageDispatcher(MessageDispatcher): - """Background worker that consumes RPC tasks from a broker, coordinating with the store.""" - - def __init__( - self, - *, - queue: MessageQueue, - supervisor: TaskSupervisor, - store: MessageStateStore, - request_runner: RequestRunner, - notification_runner: NotificationRunner, - ) -> None: - self._queue = queue - self._supervisor = supervisor - self._store = store - self._request_runner = request_runner - self._notification_runner = notification_runner - self._task: asyncio.Task[None] | None = None - - def start(self) -> None: - if self._task is not None: - msg = "dispatcher already started" - raise RuntimeError(msg) - self._task = self._supervisor.create(self._run(), name="acp.Dispatcher.loop") - - async def _run(self) -> None: - try: - async for task in self._queue: - try: - if task.kind is RpcTaskKind.REQUEST: - await self._dispatch_request(task.message) - else: - await self._dispatch_notification(task.message) - finally: - self._queue.task_done() - except asyncio.CancelledError: - return - - async def stop(self) -> None: - await self._queue.close() - if self._task is not None: - with suppress(asyncio.CancelledError): - await self._task - self._task = None - - async def _dispatch_request(self, message: dict[str, Any]) -> None: - record = self._store.begin_incoming(message.get("method", ""), message.get("params")) - - async def runner() -> None: - try: - result = await self._request_runner(message) - except Exception as exc: - self._store.fail_incoming(record, exc) - raise - else: - self._store.complete_incoming(record, result) - - self._supervisor.create(runner(), name="acp.Dispatcher.request") - - async def _dispatch_notification(self, message: dict[str, Any]) -> None: - async def runner() -> None: - await self._notification_runner(message) - - self._supervisor.create(runner(), name="acp.Dispatcher.notification") diff --git a/src/acp/task/queue.py b/src/acp/task/queue.py deleted file mode 100644 index 6052635..0000000 --- a/src/acp/task/queue.py +++ /dev/null @@ -1,67 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections.abc import AsyncIterator -from contextlib import suppress -from typing import Protocol - -from . import RpcTask - -__all__ = ["InMemoryMessageQueue", "MessageQueue"] - - -class MessageQueue(Protocol): - async def publish(self, task: RpcTask) -> None: ... - - async def close(self) -> None: ... - - def task_done(self) -> None: ... - - async def join(self) -> None: ... - - def __aiter__(self) -> AsyncIterator[RpcTask]: ... - - -class InMemoryMessageQueue: - """Simple in-memory broker for RPC task dispatch.""" - - def __init__(self, *, maxsize: int = 0) -> None: - self._queue: asyncio.Queue[RpcTask | None] = asyncio.Queue(maxsize=maxsize) - self._closed = False - - async def publish(self, task: RpcTask) -> None: - if self._closed: - msg = "mssage queue already closed" - raise RuntimeError(msg) - await self._queue.put(task) - - async def close(self) -> None: - if self._closed: - return - self._closed = True - await self._queue.put(None) - - async def join(self) -> None: - await self._queue.join() - - def task_done(self) -> None: - with suppress(ValueError): - self._queue.task_done() - - def __aiter__(self) -> AsyncIterator[RpcTask]: - return _QueueIterator(self) - - -class _QueueIterator: - def __init__(self, queue: InMemoryMessageQueue) -> None: - self._queue = queue - - def __aiter__(self) -> _QueueIterator: - return self - - async def __anext__(self) -> RpcTask: - item = await self._queue._queue.get() - if item is None: - self._queue.task_done() - raise StopAsyncIteration - return item diff --git a/src/acp/task/sender.py b/src/acp/task/sender.py index 5662af2..613535b 100644 --- a/src/acp/task/sender.py +++ b/src/acp/task/sender.py @@ -4,16 +4,12 @@ import contextlib import json import logging -from collections.abc import Callable from dataclasses import dataclass from typing import Any from .supervisor import TaskSupervisor -__all__ = ["MessageSender", "SenderFactory"] - - -SenderFactory = Callable[[asyncio.StreamWriter, TaskSupervisor], "MessageSender"] +__all__ = ["MessageSender"] @dataclass(slots=True) diff --git a/src/acp/task/state.py b/src/acp/task/state.py deleted file mode 100644 index 65baf0c..0000000 --- a/src/acp/task/state.py +++ /dev/null @@ -1,84 +0,0 @@ -from __future__ import annotations - -import asyncio -from dataclasses import dataclass -from typing import Any, Protocol - -__all__ = [ - "InMemoryMessageStateStore", - "IncomingMessage", - "MessageStateStore", - "OutgoingMessage", -] - - -@dataclass(slots=True) -class OutgoingMessage: - request_id: int - method: str - future: asyncio.Future[Any] - - -@dataclass(slots=True) -class IncomingMessage: - method: str - params: Any - status: str = "pending" - result: Any = None - error: Any = None - - -class MessageStateStore(Protocol): - def register_outgoing(self, request_id: int, method: str) -> asyncio.Future[Any]: ... - - def resolve_outgoing(self, request_id: int, result: Any) -> None: ... - - def reject_outgoing(self, request_id: int, error: Any) -> None: ... - - def reject_all_outgoing(self, error: Any) -> None: ... - - def begin_incoming(self, method: str, params: Any) -> IncomingMessage: ... - - def complete_incoming(self, record: IncomingMessage, result: Any) -> None: ... - - def fail_incoming(self, record: IncomingMessage, error: Any) -> None: ... - - -class InMemoryMessageStateStore(MessageStateStore): - def __init__(self) -> None: - self._outgoing: dict[int, OutgoingMessage] = {} - self._incoming: list[IncomingMessage] = [] - - def register_outgoing(self, request_id: int, method: str) -> asyncio.Future[Any]: - future: asyncio.Future[Any] = asyncio.get_running_loop().create_future() - self._outgoing[request_id] = OutgoingMessage(request_id, method, future) - return future - - def resolve_outgoing(self, request_id: int, result: Any) -> None: - record = self._outgoing.pop(request_id, None) - if record and not record.future.done(): - record.future.set_result(result) - - def reject_outgoing(self, request_id: int, error: Any) -> None: - record = self._outgoing.pop(request_id, None) - if record and not record.future.done(): - record.future.set_exception(error) - - def reject_all_outgoing(self, error: Any) -> None: - for record in self._outgoing.values(): - if not record.future.done(): - record.future.set_exception(error) - self._outgoing.clear() - - def begin_incoming(self, method: str, params: Any) -> IncomingMessage: - record = IncomingMessage(method=method, params=params) - self._incoming.append(record) - return record - - def complete_incoming(self, record: IncomingMessage, result: Any) -> None: - record.status = "completed" - record.result = result - - def fail_incoming(self, record: IncomingMessage, error: Any) -> None: - record.status = "failed" - record.error = error diff --git a/tests/test_connection_recovery.py b/tests/test_connection_recovery.py index 95769f6..77059d3 100644 --- a/tests/test_connection_recovery.py +++ b/tests/test_connection_recovery.py @@ -34,7 +34,7 @@ async def test_receive_loop_handles_oversized_frame(caplog: pytest.LogCaptureFix conn, reader = _make_connection(limit=128) processed: list[str] = [] - async def tracking_process(message: dict[str, Any]) -> None: + def tracking_process(message: dict[str, Any]) -> None: processed.append(message["method"]) conn._process_message = tracking_process # type: ignore[method-assign] @@ -56,7 +56,7 @@ async def test_receive_loop_handles_consecutive_oversized_frames() -> None: conn, reader = _make_connection(limit=128) processed: list[str] = [] - async def tracking_process(message: dict[str, Any]) -> None: + def tracking_process(message: dict[str, Any]) -> None: processed.append(message["method"]) conn._process_message = tracking_process # type: ignore[method-assign] diff --git a/tests/test_core.py b/tests/test_core.py index 571dd7f..8ae2219 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -10,46 +10,34 @@ @pytest.mark.asyncio -async def test_run_agent_closes_connection_when_cancelled(server, agent) -> None: - sender_created = asyncio.Event() - sender_closed = asyncio.Event() - dispatcher_started = asyncio.Event() - dispatcher_stopped = asyncio.Event() +async def test_run_agent_closes_connection_when_cancelled(agent) -> None: + receive_started = asyncio.Event() + transport_closed = asyncio.Event() - class TrackingSender: - def __init__(self, writer: asyncio.StreamWriter, supervisor: Any) -> None: - sender_created.set() + class TrackingTransport: + async def receive(self) -> dict[str, Any] | None: + receive_started.set() + await asyncio.Event().wait() + return None - async def send(self, payload: dict[str, Any]) -> None: + async def send(self, message: dict[str, Any]) -> None: msg = "test does not send messages" raise AssertionError(msg) async def close(self) -> None: - sender_closed.set() - - class TrackingDispatcher: - def start(self) -> None: - dispatcher_started.set() - - async def stop(self) -> None: - dispatcher_stopped.set() + transport_closed.set() task = asyncio.create_task( run_agent( agent, - server.server_writer, - server.server_reader, - sender_factory=TrackingSender, - dispatcher_factory=lambda *args: TrackingDispatcher(), + TrackingTransport(), ) ) - await asyncio.wait_for(sender_created.wait(), timeout=1) - await asyncio.wait_for(dispatcher_started.wait(), timeout=1) + await asyncio.wait_for(receive_started.wait(), timeout=1) task.cancel() with contextlib.suppress(asyncio.CancelledError): await asyncio.wait_for(task, timeout=1) - await asyncio.wait_for(dispatcher_stopped.wait(), timeout=1) - await asyncio.wait_for(sender_closed.wait(), timeout=1) + await asyncio.wait_for(transport_closed.wait(), timeout=1) diff --git a/tests/test_request_error_logging.py b/tests/test_request_error_logging.py index 3581804..12f3ec4 100644 --- a/tests/test_request_error_logging.py +++ b/tests/test_request_error_logging.py @@ -8,39 +8,35 @@ from __future__ import annotations -import asyncio import logging from typing import Any -from unittest.mock import MagicMock import pytest +from acp._transport import Transport from acp.connection import Connection, MethodHandler -from acp.exceptions import RequestError -class _RecordingSender: - """Duck-typed MessageSender that records outgoing frames instead of writing them.""" +class _RecordingTransport: + """Message transport that records outgoing frames.""" - def __init__(self, writer: asyncio.StreamWriter, supervisor: Any) -> None: + def __init__(self) -> None: self.sent: list[dict[str, Any]] = [] - async def send(self, payload: dict[str, Any]) -> None: - self.sent.append(payload) + async def send(self, message: dict[str, Any]) -> None: + self.sent.append(message) + + async def receive(self) -> dict[str, Any] | None: + return None async def close(self) -> None: pass -def _make_connection(handler: MethodHandler) -> tuple[Connection, _RecordingSender]: - captured: dict[str, _RecordingSender] = {} - - def sender_factory(writer: asyncio.StreamWriter, supervisor: Any) -> _RecordingSender: - captured["sender"] = _RecordingSender(writer, supervisor) - return captured["sender"] - - conn = Connection(handler, MagicMock(), MagicMock(), sender_factory=sender_factory, listening=False) - return conn, captured["sender"] +def _make_connection(handler: MethodHandler) -> tuple[Connection, _RecordingTransport]: + transport = _RecordingTransport() + conn = Connection(handler, transport, listening=False) + return conn, transport async def _raising_handler(method: str, params: Any, is_notification: bool) -> Any: @@ -62,25 +58,19 @@ def _assert_logged_runtime_error(caplog: pytest.LogCaptureFixture, method: str) @pytest.mark.asyncio -async def test_run_request_unhandled_exception_is_logged_and_returned_as_internal_error(caplog): - conn, sender = _make_connection(_raising_handler) +async def test_run_request_unhandled_exception_is_logged_and_sent_as_internal_error(caplog): + conn, transport = _make_connection(_raising_handler) request = {"jsonrpc": "2.0", "id": 7, "method": "explode", "params": None} try: - with caplog.at_level(logging.ERROR), pytest.raises(RequestError) as exc_info: + with caplog.at_level(logging.ERROR): await conn._run_request(request) finally: await conn.close() - # The handler exception is re-raised as a JSON-RPC internal error... - raised = exc_info.value - assert isinstance(raised, RequestError) - assert raised.code == -32603 - assert raised.data == {"details": "kaboom"} - - # ...and exactly one error frame carrying the handler's message is written to the peer. - assert len(sender.sent) == 1 - response = sender.sent[0] + # A handled application exception becomes exactly one JSON-RPC error frame. + assert len(transport.sent) == 1 + response = transport.sent[0] assert response["id"] == 7 assert "result" not in response assert response["error"] == {"code": -32603, "message": "Internal error", "data": {"details": "kaboom"}} @@ -91,7 +81,7 @@ async def test_run_request_unhandled_exception_is_logged_and_returned_as_interna @pytest.mark.asyncio async def test_run_notification_unhandled_exception_is_logged_and_not_answered(caplog): - conn, sender = _make_connection(_raising_handler) + conn, transport = _make_connection(_raising_handler) notification = {"jsonrpc": "2.0", "method": "session/cancel", "params": {"sessionId": "s1"}} try: @@ -102,7 +92,29 @@ async def test_run_notification_unhandled_exception_is_logged_and_not_answered(c # A notification has no response: the error is neither raised nor written to the wire. assert result is None - assert sender.sent == [] + assert transport.sent == [] # It must still be logged — previously contextlib.suppress dropped it silently. _assert_logged_runtime_error(caplog, "session/cancel") + + +@pytest.mark.asyncio +async def test_response_send_failure_is_not_mapped_to_a_handler_error(caplog): + class _FailingTransport(_RecordingTransport): + async def send(self, message: dict[str, Any]) -> None: + raise ConnectionError("send failed") + + async def successful_handler(method: str, params: Any, is_notification: bool) -> Any: + return {"ok": True} + + transport: Transport = _FailingTransport() + conn = Connection(successful_handler, transport, listening=False) + request = {"jsonrpc": "2.0", "id": 8, "method": "succeed", "params": None} + + try: + with caplog.at_level(logging.ERROR), pytest.raises(ConnectionError, match="send failed"): + await conn._run_request(request) + finally: + await conn.close() + + assert "Unhandled error while handling request method=succeed" not in caplog.text diff --git a/tests/test_rpc.py b/tests/test_rpc.py index 5e8a917..1d29465 100644 --- a/tests/test_rpc.py +++ b/tests/test_rpc.py @@ -299,6 +299,37 @@ async def test_new_requests_fail_fast_after_remote_eof(server): await conn.close() +@pytest.mark.asyncio +async def test_cancelled_request_is_removed_from_pending_map() -> None: + sent = asyncio.Event() + + class _NeverRespondTransport: + async def send(self, message: dict[str, Any]) -> None: + sent.set() + + async def receive(self) -> dict[str, Any] | None: + await asyncio.Event().wait() + return None + + async def close(self) -> None: + pass + + conn = Connection( + lambda method, params, is_notification: None, + _NeverRespondTransport(), + listening=False, + ) + request = asyncio.create_task(conn.send_request("ping")) + + await sent.wait() + request.cancel() + with pytest.raises(asyncio.CancelledError): + await request + + assert conn._pending == {} + await conn.close() + + @pytest.mark.asyncio async def test_invalid_params_results_in_error_response(connect, server): # Only start agent-side (server) so we can inject raw request from client socket