From b9c38ea6f0554cde01fbc6c9f26b079e499db3f9 Mon Sep 17 00:00:00 2001 From: Sebastiaan la Fleur Date: Wed, 23 Sep 2026 10:37:35 +0200 Subject: [PATCH 1/4] Fix async life cycle management to prevent lost references to tasks. --- src/s2python/connection/async_/connection.py | 162 +++++--- tests/unit/async_connection_test.py | 385 +++++++++++++++++++ tests/unit/reception_status_awaiter_test.py | 8 +- 3 files changed, 489 insertions(+), 66 deletions(-) create mode 100644 tests/unit/async_connection_test.py diff --git a/src/s2python/connection/async_/connection.py b/src/s2python/connection/async_/connection.py index d5c4cdd..dd1a5db 100644 --- a/src/s2python/connection/async_/connection.py +++ b/src/s2python/connection/async_/connection.py @@ -2,7 +2,7 @@ import json import logging import uuid -from typing import Optional, Type +from typing import Optional, Type, cast from s2python.connection.connection_events import ConnectionStopped from s2python.connection.async_.medium.s2_medium import ( @@ -16,8 +16,14 @@ ReceptionStatusValues, ReceptionStatus, ) -from s2python.connection.async_.message_handlers import MessageHandlers, S2EventHandlerAsync -from s2python.connection.errors import PermanentConnectionError, CouldNotReceiveStatusReceptionError +from s2python.connection.async_.message_handlers import ( + MessageHandlers, + S2EventHandlerAsync, +) +from s2python.connection.errors import ( + PermanentConnectionError, + CouldNotReceiveStatusReceptionError, +) from s2python.connection.types import S2ConnectionEventsAndMessages from s2python.reception_status_awaiter import ReceptionStatusAwaiter from s2python.s2_parser import S2Parser @@ -44,7 +50,9 @@ def __init__( medium: S2MediumConnection, eventloop: Optional[asyncio.AbstractEventLoop] = None, ) -> None: - self._eventloop = eventloop if eventloop is not None else asyncio.get_event_loop() + self._eventloop = ( + eventloop if eventloop is not None else asyncio.get_event_loop() + ) self._stop_event = asyncio.Event() self._reception_status_awaiter = ReceptionStatusAwaiter() @@ -80,36 +88,36 @@ async def run(self) -> None: "Cannot start the S2 connection if the underlying medium is closed." ) - background_tasks = [ - self._eventloop.create_task(self._receive_messages()), - self._eventloop.create_task(self._wait_till_stop()), - self._eventloop.create_task(self._handle_received_messages()), - ] - - await self._handlers.handle_event(self, ConnectionStarted()) - - (done, pending) = await asyncio.wait(background_tasks, return_when=asyncio.FIRST_COMPLETED) - - await self._handlers.handle_event(self, ConnectionStopped()) - - for task in pending: - try: - task.cancel() - await task - except (asyncio.CancelledError, Exception): # pylint: disable=broad-exception-caught - pass + background_tasks = [] + try: + background_tasks.append( + self._eventloop.create_task(self._receive_messages()) + ) + background_tasks.append(self._eventloop.create_task(self._wait_till_stop())) + background_tasks.append( + self._eventloop.create_task(self._handle_received_messages()) + ) - for task in done: - try: - await task - except asyncio.CancelledError: - pass - except MediumClosedConnectionError: - logger.info("The other party closed the websocket connection.") - except Exception: # pylint: disable=broad-exception-caught - logger.exception( - "An error occurred in the S2 connection. Terminating current connection." - ) + await self._handlers.handle_event(self, ConnectionStarted()) + await asyncio.wait(background_tasks, return_when=asyncio.FIRST_COMPLETED) + await self._handlers.handle_event(self, ConnectionStopped()) + finally: + for task in background_tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*background_tasks, return_exceptions=True) + + for task in background_tasks: + try: + task.result() + except asyncio.CancelledError: + pass + except MediumClosedConnectionError: + logger.info("The other party closed the websocket connection.") + except Exception: # pylint: disable=broad-exception-caught + logger.exception( + "An error occurred in the S2 connection. Terminating current connection." + ) async def _handle_received_messages(self) -> None: while not self._stop_event.is_set(): @@ -131,7 +139,9 @@ async def _receive_messages(self) -> None: except json.JSONDecodeError: await self.send_and_forget( ReceptionStatus( - subject_message_id=uuid.UUID("00000000-0000-0000-0000-000000000000"), + subject_message_id=uuid.UUID( + "00000000-0000-0000-0000-000000000000" + ), status=ReceptionStatusValues.INVALID_DATA, diagnostic_label="Not valid json.", ) @@ -147,7 +157,9 @@ async def _receive_messages(self) -> None: ) else: await self.respond_with_reception_status( - subject_message_id=uuid.UUID("00000000-0000-0000-0000-000000000000"), + subject_message_id=uuid.UUID( + "00000000-0000-0000-0000-000000000000" + ), status=ReceptionStatusValues.INVALID_DATA, diagnostic_label="Message appears valid json but could not find a message_id field.", ) @@ -159,7 +171,9 @@ async def _receive_messages(self) -> None: "Message is a reception status for %s so registering in cache.", s2_msg.subject_message_id, ) - await self._reception_status_awaiter.receive_reception_status(s2_msg) + await self._reception_status_awaiter.receive_reception_status( + s2_msg + ) else: logger.debug( "Message is not a reception status, putting it in the received messages queue." @@ -167,7 +181,9 @@ async def _receive_messages(self) -> None: await self._received_messages.put(s2_msg) def register_handler( - self, event_type: Type[S2ConnectionEventsAndMessages], handler: S2EventHandlerAsync + self, + event_type: Type[S2ConnectionEventsAndMessages], + handler: S2EventHandlerAsync, ) -> None: """Register a handler for a specific S2 message type. @@ -176,7 +192,9 @@ def register_handler( """ self._handlers.register_handler(event_type, handler) - def unregister_handler(self, s2_message_type: Type[S2ConnectionEventsAndMessages]) -> None: + def unregister_handler( + self, s2_message_type: Type[S2ConnectionEventsAndMessages] + ) -> None: self._handlers.unregister_handler(s2_message_type) async def send_and_forget(self, s2_msg: S2Message) -> None: @@ -189,9 +207,14 @@ async def send_and_forget(self, s2_msg: S2Message) -> None: raise async def respond_with_reception_status( - self, subject_message_id: uuid.UUID, status: ReceptionStatusValues, diagnostic_label: str + self, + subject_message_id: uuid.UUID, + status: ReceptionStatusValues, + diagnostic_label: str, ) -> None: - logger.debug("Responding to message %s with status %s", subject_message_id, status) + logger.debug( + "Responding to message %s with status %s", subject_message_id, status + ) await self.send_and_forget( ReceptionStatus( subject_message_id=subject_message_id, @@ -226,34 +249,45 @@ async def send_msg_and_await_reception_status( s2_msg.message_id, timeout_reception_status, ) - reception_status_task = self._eventloop.create_task( - self._reception_status_awaiter.wait_for_reception_status( - s2_msg.message_id, timeout_reception_status + tasks: list[asyncio.Task] = [] + done: set[asyncio.Task] = set() + try: + reception_status_task = self._eventloop.create_task( + self._reception_status_awaiter.wait_for_reception_status( + s2_msg.message_id, timeout_reception_status + ) ) - ) - stop_event_task = self._eventloop.create_task(self._wait_till_stop()) - - (done, pending) = await asyncio.wait( - [reception_status_task, stop_event_task], return_when=asyncio.FIRST_COMPLETED - ) - - for task in pending: - try: - task.cancel() - await task - except (asyncio.CancelledError, Exception): # pylint: disable=broad-exception-caught - pass + tasks.append(reception_status_task) + stop_event_task = self._eventloop.create_task(self._wait_till_stop()) + tasks.append(stop_event_task) + + (done, _) = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED) + finally: + for task in tasks: + if not task.done(): + task.cancel() + # Collect every task result so concurrently completed tasks cannot leak exceptions. + results = await asyncio.gather(*tasks, return_exceptions=True) + + reception_status_result, stop_event_result = results + if reception_status_task in done and isinstance( + reception_status_result, BaseException + ): + if isinstance( + reception_status_result, (TimeoutError, asyncio.TimeoutError) + ): + logger.error( + "Did not receive a reception status on time for %s", + s2_msg.message_id, + ) + self._stop_event.set() + raise reception_status_result + if stop_event_task in done and isinstance(stop_event_result, BaseException): + raise stop_event_result if reception_status_task in done: - try: - reception_status = await reception_status_task - except (TimeoutError, asyncio.TimeoutError): - logger.error("Did not receive a reception status on time for %s", s2_msg.message_id) - self._stop_event.set() - raise + reception_status = cast(ReceptionStatus, reception_status_result) else: - # stop_event_task in done - await stop_event_task raise CouldNotReceiveStatusReceptionError( f"Connection stopped while waiting for ReceptionStatus for message {s2_msg.message_id}" ) diff --git a/tests/unit/async_connection_test.py b/tests/unit/async_connection_test.py new file mode 100644 index 0000000..b47e88b --- /dev/null +++ b/tests/unit/async_connection_test.py @@ -0,0 +1,385 @@ +"""Tests for async connection task management.""" + +import asyncio +import uuid +from unittest import IsolatedAsyncioTestCase +from unittest.mock import AsyncMock, Mock + +from s2python.common import ReceptionStatus, ReceptionStatusValues +from s2python.connection.async_.connection import S2AsyncConnection +from s2python.connection.async_.medium.s2_medium import S2AsyncMediumConnection +from s2python.connection.connection_events import ConnectionStarted, ConnectionStopped +from s2python.connection.errors import ( + CouldNotReceiveStatusReceptionError, + PermanentConnectionError, +) + + +class _Medium(S2AsyncMediumConnection): + async def is_connected(self) -> bool: + return True + + async def messages(self): + if False: + yield "" + + async def send(self, message: str) -> None: + pass + + +class _BlockingReceptionStatusAwaiter: + def __init__(self) -> None: + self.waiting = asyncio.Event() + self.done = asyncio.Event() + + async def wait_for_reception_status( + self, message_id: uuid.UUID, timeout_reception_status: float + ) -> ReceptionStatus: # pyright: ignore [reportReturnType] + self.waiting.set() + try: + await asyncio.Event().wait() + finally: + self.done.set() + + +class _ResultReceptionStatusAwaiter: + def __init__(self, result) -> None: + self.result = result + + async def wait_for_reception_status( + self, message_id: uuid.UUID, timeout_reception_status: float + ) -> ReceptionStatus: + if isinstance(self.result, BaseException): + raise self.result + return self.result + + +class _FailingStopConnection(S2AsyncConnection): + async def _wait_till_stop(self) -> None: + raise RuntimeError("stop waiter failed") + + +class _BlockingStopConnection(S2AsyncConnection): + def __init__(self, medium: S2AsyncMediumConnection) -> None: + super().__init__(medium) + self.stop_waiting = asyncio.Event() + self.stop_waiter_done = asyncio.Event() + + async def _wait_till_stop(self) -> None: + self.stop_waiting.set() + try: + await super()._wait_till_stop() + finally: + self.stop_waiter_done.set() + + +class _RecordingTaskFactory: + def __init__(self) -> None: + self.loop = asyncio.get_event_loop() + self.tasks = [] + + def create_task(self, coroutine): + task = self.loop.create_task(coroutine) + self.tasks.append(task) + return task + + +class _LifecycleConnection(S2AsyncConnection): + def __init__(self, task_factory: _RecordingTaskFactory) -> None: + super().__init__(_Medium(), eventloop=task_factory) # type: ignore[arg-type] + self.receive_started = asyncio.Event() + self.handle_started = asyncio.Event() + self.receive_done = asyncio.Event() + self.handle_done = asyncio.Event() + + async def _receive_messages(self) -> None: + self.receive_started.set() + try: + await asyncio.Event().wait() + finally: + self.receive_done.set() + + async def _handle_received_messages(self) -> None: + self.handle_started.set() + try: + await asyncio.Event().wait() + finally: + self.handle_done.set() + + +class _FailingReceiveConnection(_LifecycleConnection): + async def _receive_messages(self) -> None: + self.receive_started.set() + self.receive_done.set() + raise RuntimeError("receive failed") + + +class AsyncConnectionTest(IsolatedAsyncioTestCase): + def setUp(self) -> None: + self.message = Mock() + self.message.message_id = uuid.uuid4() + self.message.to_json.return_value = "{}" + + async def wait_for_background_tasks(self, connection: _LifecycleConnection) -> None: + await connection.receive_started.wait() + await connection.handle_started.wait() + + def reception_status( + self, status: ReceptionStatusValues = ReceptionStatusValues.OK + ): + return ReceptionStatus( # pyright: ignore[reportCallIssue] + subject_message_id=self.message.message_id, status=status + ) + + async def test__send_msg_and_await_reception_status__send_msg_cancellation_drains_reception_status_task( + self, + ) -> None: + # Arrange + connection = S2AsyncConnection(_Medium()) + awaiter = _BlockingReceptionStatusAwaiter() + connection._reception_status_awaiter = awaiter + + send_task = asyncio.create_task( + connection.send_msg_and_await_reception_status(self.message) + ) + await awaiter.waiting.wait() + + # Act + send_task.cancel() + + # Assert + with self.assertRaises(asyncio.CancelledError): + await send_task + + self.assertTrue(awaiter.done.is_set()) + + async def test__send_msg_and_await_reception_status__returns_reception_status( + self, + ) -> None: + # Arrange + connection = _BlockingStopConnection(_Medium()) + expected_status = self.reception_status() + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + expected_status + ) + + # Act + received_status = await connection.send_msg_and_await_reception_status( + self.message + ) + + # Assert + self.assertEqual(expected_status, received_status) + self.assertTrue(connection.stop_waiter_done.is_set()) + + async def test__send_msg_and_await_reception_status__times_out_and_stops_connection( + self, + ) -> None: + # Arrange + connection = _BlockingStopConnection(_Medium()) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + asyncio.TimeoutError() + ) + + # Act & Assert + with self.assertRaises(asyncio.TimeoutError): + await connection.send_msg_and_await_reception_status(self.message) + + self.assertTrue(connection._stop_event.is_set()) + self.assertTrue(connection.stop_waiter_done.is_set()) + + async def test__send_msg_and_await_reception_status__real_timeout_stops_connection( + self, + ) -> None: + connection = S2AsyncConnection(_Medium()) + + with self.assertRaises(asyncio.TimeoutError): + await connection.send_msg_and_await_reception_status( + self.message, timeout_reception_status=0 + ) + + self.assertTrue(connection._stop_event.is_set()) + self.assertEqual({}, connection._reception_status_awaiter.awaiting) + + async def test__send_msg_and_await_reception_status__stopping_connection_drains_reception_status_task( + self, + ) -> None: + # Arrange + connection = S2AsyncConnection(_Medium()) + awaiter = _BlockingReceptionStatusAwaiter() + connection._reception_status_awaiter = awaiter + send_task = asyncio.create_task( + connection.send_msg_and_await_reception_status(self.message) + ) + await awaiter.waiting.wait() + + # Act + await connection.stop() + + # Assert + with self.assertRaisesRegex( + CouldNotReceiveStatusReceptionError, + "Connection stopped while waiting for ReceptionStatus", + ): + await send_task + self.assertTrue(awaiter.done.is_set()) + + async def test__send_msg_and_await_reception_status__send_msg_observes_stop_task_exception( + self, + ) -> None: + # Arrange + connection = _FailingStopConnection(_Medium()) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + self.reception_status() + ) + + # Act & Assert + with self.assertRaisesRegex(RuntimeError, "stop waiter failed"): + await connection.send_msg_and_await_reception_status(self.message) + + async def test__send_msg_and_await_reception_status__propagates_reception_status_exception( + self, + ) -> None: + # Arrange + connection = _BlockingStopConnection(_Medium()) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + ValueError("reception status waiter failed") + ) + + # Act & Assert + with self.assertRaisesRegex(ValueError, "reception status waiter failed"): + await connection.send_msg_and_await_reception_status(self.message) + + self.assertTrue(connection.stop_waiter_done.is_set()) + + async def test__send_msg_and_await_reception_status__propagates_child_cancellation( + self, + ) -> None: + # Arrange + connection = _BlockingStopConnection(_Medium()) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + asyncio.CancelledError() + ) + + # Act & Assert + with self.assertRaises(asyncio.CancelledError): + await connection.send_msg_and_await_reception_status(self.message) + + self.assertTrue(connection.stop_waiter_done.is_set()) + + async def test__send_msg_and_await_reception_status__handles_error_statuses( + self, + ) -> None: + # Arrange + connection = S2AsyncConnection(_Medium()) + permanent_error_status = self.reception_status( + ReceptionStatusValues.PERMANENT_ERROR + ) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + permanent_error_status + ) + + # Act & Assert + with self.assertRaises(PermanentConnectionError): + await connection.send_msg_and_await_reception_status(self.message) + + connection = S2AsyncConnection(_Medium()) + connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( + permanent_error_status + ) + received_status = await connection.send_msg_and_await_reception_status( + self.message, raise_on_error=False + ) + + self.assertEqual(permanent_error_status, received_status) + + async def test__run__graceful_stop_drains_background_tasks(self) -> None: + # Arrange + task_factory = _RecordingTaskFactory() + connection = _LifecycleConnection(task_factory) + connection._handlers.handle_event = AsyncMock() + run_task = asyncio.create_task(connection.run()) + await self.wait_for_background_tasks(connection) + + # Act + await connection.stop() + await run_task + + # Assert + self.assertTrue(all(task.done() for task in task_factory.tasks)) + self.assertTrue(connection.receive_done.is_set()) + self.assertTrue(connection.handle_done.is_set()) + handled_events = [ + type(call.args[1]) + for call in connection._handlers.handle_event.await_args_list + ] + self.assertEqual([ConnectionStarted, ConnectionStopped], handled_events) + + async def test__run__parent_cancellation_drains_background_tasks(self) -> None: + # Arrange + task_factory = _RecordingTaskFactory() + connection = _LifecycleConnection(task_factory) + connection._handlers.handle_event = AsyncMock() + run_task = asyncio.create_task(connection.run()) + await self.wait_for_background_tasks(connection) + + # Act + run_task.cancel() + + # Assert + with self.assertRaises(asyncio.CancelledError): + await run_task + self.assertTrue(all(task.done() for task in task_factory.tasks)) + self.assertTrue(connection.receive_done.is_set()) + self.assertTrue(connection.handle_done.is_set()) + + async def test__run__started_handler_failure_drains_background_tasks(self) -> None: + # Arrange + task_factory = _RecordingTaskFactory() + connection = _LifecycleConnection(task_factory) + connection._handlers.handle_event = AsyncMock( + side_effect=RuntimeError("handler failed") + ) + + # Act & Assert + with self.assertRaisesRegex(RuntimeError, "handler failed"): + await connection.run() + self.assertTrue(all(task.done() for task in task_factory.tasks)) + + async def test__run__stopped_handler_failure_drains_background_tasks(self) -> None: + # Arrange + task_factory = _RecordingTaskFactory() + connection = _LifecycleConnection(task_factory) + + async def handle_event(_, event) -> None: + if isinstance(event, ConnectionStopped): + raise RuntimeError("handler failed") + + connection._handlers.handle_event = handle_event + run_task = asyncio.create_task(connection.run()) + await self.wait_for_background_tasks(connection) + + # Act + await connection.stop() + + # Assert + with self.assertRaisesRegex(RuntimeError, "handler failed"): + await run_task + self.assertTrue(all(task.done() for task in task_factory.tasks)) + self.assertTrue(connection.receive_done.is_set()) + self.assertTrue(connection.handle_done.is_set()) + + async def test__run__background_failure_drains_remaining_tasks(self) -> None: + # Arrange + task_factory = _RecordingTaskFactory() + connection = _FailingReceiveConnection(task_factory) + connection._handlers.handle_event = AsyncMock() + + # Act + with self.assertLogs("s2python", level="ERROR") as logs: + await connection.run() + + # Assert + self.assertTrue(all(task.done() for task in task_factory.tasks)) + self.assertTrue(connection.handle_done.is_set()) + self.assertTrue(any("receive failed" in message for message in logs.output)) diff --git a/tests/unit/reception_status_awaiter_test.py b/tests/unit/reception_status_awaiter_test.py index 5e27c83..1d83ce5 100644 --- a/tests/unit/reception_status_awaiter_test.py +++ b/tests/unit/reception_status_awaiter_test.py @@ -93,8 +93,12 @@ async def test__wait_for_reception_status__multiple_receive_while_waiting(self): self.assertTrue(should_be_waiting_still_1) self.assertTrue(should_be_waiting_still_2) - successful_results = [result for result in results if not isinstance(result, Exception)] - exception_results = [result for result in results if isinstance(result, Exception)] + successful_results = [ + result for result in results if not isinstance(result, Exception) + ] + exception_results = [ + result for result in results if isinstance(result, Exception) + ] self.assertEqual(1, len(successful_results)) self.assertEqual(1, len(exception_results)) From 12f930089bc6795e3010a43bf9d3f90c48bbf293 Mon Sep 17 00:00:00 2001 From: Sebastiaan la Fleur Date: Wed, 23 Sep 2026 11:24:33 +0200 Subject: [PATCH 2/4] Simplify some of the life cycle management in send_msg_and_await_reception_status --- src/s2python/connection/async_/connection.py | 31 +++++++++----------- 1 file changed, 14 insertions(+), 17 deletions(-) diff --git a/src/s2python/connection/async_/connection.py b/src/s2python/connection/async_/connection.py index dd1a5db..3161fe6 100644 --- a/src/s2python/connection/async_/connection.py +++ b/src/s2python/connection/async_/connection.py @@ -2,7 +2,7 @@ import json import logging import uuid -from typing import Optional, Type, cast +from typing import Optional, Type from s2python.connection.connection_events import ConnectionStopped from s2python.connection.async_.medium.s2_medium import ( @@ -250,7 +250,7 @@ async def send_msg_and_await_reception_status( timeout_reception_status, ) tasks: list[asyncio.Task] = [] - done: set[asyncio.Task] = set() + done: set[asyncio.Task] try: reception_status_task = self._eventloop.create_task( self._reception_status_awaiter.wait_for_reception_status( @@ -267,27 +267,24 @@ async def send_msg_and_await_reception_status( if not task.done(): task.cancel() # Collect every task result so concurrently completed tasks cannot leak exceptions. - results = await asyncio.gather(*tasks, return_exceptions=True) - - reception_status_result, stop_event_result = results - if reception_status_task in done and isinstance( - reception_status_result, BaseException - ): - if isinstance( - reception_status_result, (TimeoutError, asyncio.TimeoutError) - ): + await asyncio.gather(*tasks, return_exceptions=True) + + reception_status = None + if reception_status_task in done: + try: + reception_status = reception_status_task.result() + except (TimeoutError, asyncio.TimeoutError): logger.error( "Did not receive a reception status on time for %s", s2_msg.message_id, ) self._stop_event.set() - raise reception_status_result - if stop_event_task in done and isinstance(stop_event_result, BaseException): - raise stop_event_result + raise - if reception_status_task in done: - reception_status = cast(ReceptionStatus, reception_status_result) - else: + if stop_event_task in done: + stop_event_task.result() + + if reception_status is None: raise CouldNotReceiveStatusReceptionError( f"Connection stopped while waiting for ReceptionStatus for message {s2_msg.message_id}" ) From 47a6091e56a74cb4ced989fe19480a0ec8d82ce9 Mon Sep 17 00:00:00 2001 From: Sebastiaan la Fleur Date: Wed, 23 Sep 2026 12:22:28 +0200 Subject: [PATCH 3/4] Allow a number of linting and typing violations under tests/ --- .pylintrc | 2 ++ ci/lint.sh | 4 +++- ci/setup_dev_environment.sh | 2 +- ci/test_unit.sh | 2 +- ci/typecheck.sh | 6 ++++-- mypy.ini | 1 + pyrightconfig.json | 7 +++++++ tests/unit/connection/__init__.py | 0 tests/unit/connection/async_/__init__.py | 0 .../async_/connection_test.py} | 18 +++++++++++++++--- tox.ini | 5 ++++- 11 files changed, 38 insertions(+), 9 deletions(-) create mode 100644 tests/unit/connection/__init__.py create mode 100644 tests/unit/connection/async_/__init__.py rename tests/unit/{async_connection_test.py => connection/async_/connection_test.py} (95%) diff --git a/.pylintrc b/.pylintrc index 0c2006b..39cefb4 100644 --- a/.pylintrc +++ b/.pylintrc @@ -11,3 +11,5 @@ ignore-paths=src/s2python/generated/ jobs=1 disable=missing-class-docstring,missing-module-docstring,too-few-public-methods,missing-function-docstring,no-member,unsubscriptable-object,line-too-long,duplicate-code +# Lint entry points disable these separately for tests; keep them enabled elsewhere. +enable=protected-access,invalid-overridden-method diff --git a/ci/lint.sh b/ci/lint.sh index c405891..20b4f6e 100755 --- a/ci/lint.sh +++ b/ci/lint.sh @@ -1,4 +1,6 @@ #!/usr/bin/env sh +set -e . .venv/bin/activate -pylint src/ tests/unit/ examples/ +pylint src/ examples/ +PYTHONPATH="src:$PYTHONPATH" pylint --disable=protected-access,invalid-overridden-method tests/unit/ diff --git a/ci/setup_dev_environment.sh b/ci/setup_dev_environment.sh index 7e10b51..98afff7 100755 --- a/ci/setup_dev_environment.sh +++ b/ci/setup_dev_environment.sh @@ -1,5 +1,5 @@ #!/bin/bash -python3.8 -m venv ./.venv/ +python3.9 -m venv ./.venv/ . ./.venv/bin/activate pip install pip-tools diff --git a/ci/test_unit.sh b/ci/test_unit.sh index 492b1ab..9a13e52 100755 --- a/ci/test_unit.sh +++ b/ci/test_unit.sh @@ -1,4 +1,4 @@ #!/usr/bin/env sh . .venv/bin/activate -PYTHONPATH="$PYTHONPATH:src/" pytest --cov=s2python --cov-report=html:./unit_test_coverage/ -v tests/unit/ +PYTHONPATH="$PYTHONPATH:src/" pytest --cov=s2python --cov-report=html:./unit_test_coverage/ -v tests/unit/ $@ diff --git a/ci/typecheck.sh b/ci/typecheck.sh index 6864b6a..bc0dc41 100755 --- a/ci/typecheck.sh +++ b/ci/typecheck.sh @@ -1,5 +1,7 @@ #!/usr/bin/env sh . .venv/bin/activate -mypy --config-file mypy.ini src/ ./tests/unit/ examples/ -pyright +status=0 +mypy --config-file mypy.ini src/ ./tests/unit/ examples/ || status=1 +pyright || status=1 +exit "$status" diff --git a/mypy.ini b/mypy.ini index 1d4ec90..e2e2ce8 100644 --- a/mypy.ini +++ b/mypy.ini @@ -12,3 +12,4 @@ disallow_incomplete_defs = false [mypy-unit.*] check_untyped_defs = true +disable_error_code = method-assign, assignment, return diff --git a/pyrightconfig.json b/pyrightconfig.json index d5e46e5..1c0c923 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -8,6 +8,13 @@ "src/s2python/generated/" ], + "executionEnvironments": [ + { + "root": "tests", + "reportAttributeAccessIssue": "none" + } + ], + "defineConstant": { "DEBUG": true } diff --git a/tests/unit/connection/__init__.py b/tests/unit/connection/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/connection/async_/__init__.py b/tests/unit/connection/async_/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/unit/async_connection_test.py b/tests/unit/connection/async_/connection_test.py similarity index 95% rename from tests/unit/async_connection_test.py rename to tests/unit/connection/async_/connection_test.py index b47e88b..46cafc1 100644 --- a/tests/unit/async_connection_test.py +++ b/tests/unit/connection/async_/connection_test.py @@ -1,6 +1,7 @@ """Tests for async connection task management.""" import asyncio +from typing import Coroutine import uuid from unittest import IsolatedAsyncioTestCase from unittest.mock import AsyncMock, Mock @@ -76,9 +77,9 @@ async def _wait_till_stop(self) -> None: class _RecordingTaskFactory: def __init__(self) -> None: self.loop = asyncio.get_event_loop() - self.tasks = [] + self.tasks: list[asyncio.Task] = [] - def create_task(self, coroutine): + def create_task(self, coroutine: Coroutine) -> asyncio.Task: task = self.loop.create_task(coroutine) self.tasks.append(task) return task @@ -171,6 +172,7 @@ async def test__send_msg_and_await_reception_status__returns_reception_status( # Assert self.assertEqual(expected_status, received_status) self.assertTrue(connection.stop_waiter_done.is_set()) + self.assertFalse(connection._stop_event.is_set()) async def test__send_msg_and_await_reception_status__times_out_and_stops_connection( self, @@ -267,7 +269,7 @@ async def test__send_msg_and_await_reception_status__propagates_child_cancellati self.assertTrue(connection.stop_waiter_done.is_set()) - async def test__send_msg_and_await_reception_status__handles_error_statuses( + async def test__send_msg_and_await_reception_status__raises_on_permanent_error( self, ) -> None: # Arrange @@ -283,14 +285,24 @@ async def test__send_msg_and_await_reception_status__handles_error_statuses( with self.assertRaises(PermanentConnectionError): await connection.send_msg_and_await_reception_status(self.message) + async def test__send_msg_and_await_reception_status__returns_permanent_error_when_raise_on_error_is_false( + self, + ) -> None: + # Arrange connection = S2AsyncConnection(_Medium()) + permanent_error_status = self.reception_status( + ReceptionStatusValues.PERMANENT_ERROR + ) connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( permanent_error_status ) + + # Act received_status = await connection.send_msg_and_await_reception_status( self.message, raise_on_error=False ) + # Assert self.assertEqual(permanent_error_status, received_status) async def test__run__graceful_stop_drains_background_tasks(self) -> None: diff --git a/tox.ini b/tox.ini index fbef9e5..ea5945a 100644 --- a/tox.ini +++ b/tox.ini @@ -39,10 +39,13 @@ commands = description = Lint the source code using pylint. skip_install = True changedir = {toxinidir} +setenv = + PYTHONPATH = {toxinidir}/src deps = -r dev-requirements.txt commands = - pylint src/ tests/unit/ + pylint src/ + pylint --disable=protected-access,invalid-overridden-method tests/unit/ [testenv:typecheck] description = Typecheck the source code using mypy. From 44d01ec1cb2db339a7340f3d931bf78bfd9806d1 Mon Sep 17 00:00:00 2001 From: Sebastiaan la Fleur Date: Wed, 23 Sep 2026 12:31:07 +0200 Subject: [PATCH 4/4] Unit test tweaks for linter and typechecker. --- .../unit/connection/async_/connection_test.py | 42 ++++++++++--------- 1 file changed, 23 insertions(+), 19 deletions(-) diff --git a/tests/unit/connection/async_/connection_test.py b/tests/unit/connection/async_/connection_test.py index 46cafc1..83d92e5 100644 --- a/tests/unit/connection/async_/connection_test.py +++ b/tests/unit/connection/async_/connection_test.py @@ -1,14 +1,17 @@ """Tests for async connection task management.""" import asyncio -from typing import Coroutine +from typing import AsyncGenerator, Coroutine import uuid from unittest import IsolatedAsyncioTestCase from unittest.mock import AsyncMock, Mock from s2python.common import ReceptionStatus, ReceptionStatusValues from s2python.connection.async_.connection import S2AsyncConnection -from s2python.connection.async_.medium.s2_medium import S2AsyncMediumConnection +from s2python.connection.async_.medium.s2_medium import ( + S2AsyncMediumConnection, + UnparsedMediumData, +) from s2python.connection.connection_events import ConnectionStarted, ConnectionStopped from s2python.connection.errors import ( CouldNotReceiveStatusReceptionError, @@ -16,13 +19,14 @@ ) -class _Medium(S2AsyncMediumConnection): +class _EmptyMessageAsyncMedium(S2AsyncMediumConnection): async def is_connected(self) -> bool: return True - async def messages(self): - if False: - yield "" + async def messages(self) -> AsyncGenerator[UnparsedMediumData, None]: + empty_messages: tuple[UnparsedMediumData, ...] = () + for message in empty_messages: + yield message async def send(self, message: str) -> None: pass @@ -34,7 +38,7 @@ def __init__(self) -> None: self.done = asyncio.Event() async def wait_for_reception_status( - self, message_id: uuid.UUID, timeout_reception_status: float + self, _: uuid.UUID, __: float ) -> ReceptionStatus: # pyright: ignore [reportReturnType] self.waiting.set() try: @@ -48,7 +52,7 @@ def __init__(self, result) -> None: self.result = result async def wait_for_reception_status( - self, message_id: uuid.UUID, timeout_reception_status: float + self, _: uuid.UUID, __: float ) -> ReceptionStatus: if isinstance(self.result, BaseException): raise self.result @@ -87,7 +91,7 @@ def create_task(self, coroutine: Coroutine) -> asyncio.Task: class _LifecycleConnection(S2AsyncConnection): def __init__(self, task_factory: _RecordingTaskFactory) -> None: - super().__init__(_Medium(), eventloop=task_factory) # type: ignore[arg-type] + super().__init__(_EmptyMessageAsyncMedium(), eventloop=task_factory) # type: ignore[arg-type] self.receive_started = asyncio.Event() self.handle_started = asyncio.Event() self.receive_done = asyncio.Event() @@ -136,7 +140,7 @@ async def test__send_msg_and_await_reception_status__send_msg_cancellation_drain self, ) -> None: # Arrange - connection = S2AsyncConnection(_Medium()) + connection = S2AsyncConnection(_EmptyMessageAsyncMedium()) awaiter = _BlockingReceptionStatusAwaiter() connection._reception_status_awaiter = awaiter @@ -158,7 +162,7 @@ async def test__send_msg_and_await_reception_status__returns_reception_status( self, ) -> None: # Arrange - connection = _BlockingStopConnection(_Medium()) + connection = _BlockingStopConnection(_EmptyMessageAsyncMedium()) expected_status = self.reception_status() connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( expected_status @@ -178,7 +182,7 @@ async def test__send_msg_and_await_reception_status__times_out_and_stops_connect self, ) -> None: # Arrange - connection = _BlockingStopConnection(_Medium()) + connection = _BlockingStopConnection(_EmptyMessageAsyncMedium()) connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( asyncio.TimeoutError() ) @@ -193,7 +197,7 @@ async def test__send_msg_and_await_reception_status__times_out_and_stops_connect async def test__send_msg_and_await_reception_status__real_timeout_stops_connection( self, ) -> None: - connection = S2AsyncConnection(_Medium()) + connection = S2AsyncConnection(_EmptyMessageAsyncMedium()) with self.assertRaises(asyncio.TimeoutError): await connection.send_msg_and_await_reception_status( @@ -207,7 +211,7 @@ async def test__send_msg_and_await_reception_status__stopping_connection_drains_ self, ) -> None: # Arrange - connection = S2AsyncConnection(_Medium()) + connection = S2AsyncConnection(_EmptyMessageAsyncMedium()) awaiter = _BlockingReceptionStatusAwaiter() connection._reception_status_awaiter = awaiter send_task = asyncio.create_task( @@ -230,7 +234,7 @@ async def test__send_msg_and_await_reception_status__send_msg_observes_stop_task self, ) -> None: # Arrange - connection = _FailingStopConnection(_Medium()) + connection = _FailingStopConnection(_EmptyMessageAsyncMedium()) connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( self.reception_status() ) @@ -243,7 +247,7 @@ async def test__send_msg_and_await_reception_status__propagates_reception_status self, ) -> None: # Arrange - connection = _BlockingStopConnection(_Medium()) + connection = _BlockingStopConnection(_EmptyMessageAsyncMedium()) connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( ValueError("reception status waiter failed") ) @@ -258,7 +262,7 @@ async def test__send_msg_and_await_reception_status__propagates_child_cancellati self, ) -> None: # Arrange - connection = _BlockingStopConnection(_Medium()) + connection = _BlockingStopConnection(_EmptyMessageAsyncMedium()) connection._reception_status_awaiter = _ResultReceptionStatusAwaiter( asyncio.CancelledError() ) @@ -273,7 +277,7 @@ async def test__send_msg_and_await_reception_status__raises_on_permanent_error( self, ) -> None: # Arrange - connection = S2AsyncConnection(_Medium()) + connection = S2AsyncConnection(_EmptyMessageAsyncMedium()) permanent_error_status = self.reception_status( ReceptionStatusValues.PERMANENT_ERROR ) @@ -289,7 +293,7 @@ async def test__send_msg_and_await_reception_status__returns_permanent_error_whe self, ) -> None: # Arrange - connection = S2AsyncConnection(_Medium()) + connection = S2AsyncConnection(_EmptyMessageAsyncMedium()) permanent_error_status = self.reception_status( ReceptionStatusValues.PERMANENT_ERROR )