diff --git a/can/notifier.py b/can/notifier.py index fd21a0662..97d762bf6 100644 --- a/can/notifier.py +++ b/can/notifier.py @@ -223,7 +223,7 @@ def _rx_thread(self, bus: BusABC) -> None: if self._loop: handle_message: Callable[[Message], Any] = functools.partial( self._loop.call_soon_threadsafe, - self._on_message_received, # type: ignore[arg-type] + self._on_message_received_with_error_handling, # type: ignore[arg-type] ) else: handle_message = self._on_message_received @@ -248,7 +248,16 @@ def _rx_thread(self, bus: BusABC) -> None: def _on_message_available(self, bus: BusABC) -> None: if msg := bus.recv(0): + self._on_message_received_with_error_handling(msg) + + def _on_message_received_with_error_handling(self, msg: Message) -> None: + try: self._on_message_received(msg) + except Exception as exc: # pylint: disable=broad-except + self.exception = exc + if not self._on_error(exc): + raise + logger.debug("suppressed exception: %s", exc) def _on_message_received(self, msg: Message) -> None: for callback in self.listeners: @@ -257,7 +266,23 @@ def _on_message_received(self, msg: Message) -> None: # Schedule coroutine and keep a reference to the task task = self._loop.create_task(res) self._tasks.add(task) - task.add_done_callback(self._tasks.discard) + task.add_done_callback(self._on_task_done) + + def _on_task_done(self, task: asyncio.Task) -> None: + self._tasks.discard(task) + if task.cancelled(): + return + + exc = task.exception() + if exc is None: + return + if not isinstance(exc, Exception): + raise exc + + self.exception = exc + if not self._on_error(exc): + raise exc + logger.debug("suppressed exception: %s", exc) def _on_error(self, exc: Exception) -> bool: """Calls ``on_error()`` for all listeners if they implement it. diff --git a/test/notifier_test.py b/test/notifier_test.py index d8512a00b..a190ab6e9 100644 --- a/test/notifier_test.py +++ b/test/notifier_test.py @@ -7,6 +7,24 @@ import can +class RaisingListener(can.Listener): + def on_message_received(self, msg: can.Message) -> None: + raise ValueError("listener failed") + + +class ErrorCollector(can.Listener): + def __init__(self, event: asyncio.Event) -> None: + self.event = event + self.errors: list[Exception] = [] + + def on_message_received(self, msg: can.Message) -> None: + pass + + def on_error(self, exc: Exception) -> None: + self.errors.append(exc) + self.event.set() + + class NotifierTest(unittest.TestCase): def test_single_bus(self): with can.Bus("test", interface="virtual", receive_own_messages=True) as bus: @@ -88,6 +106,58 @@ async def run_it(): asyncio.run(run_it()) + def test_sync_listener_error_calls_on_error_with_loop(self): + async def run_it(): + event = asyncio.Event() + collector = ErrorCollector(event) + with can.Bus( + "sync-listener-error", interface="virtual", receive_own_messages=True + ) as bus: + notifier = can.Notifier( + bus, + [RaisingListener(), collector], + 0.1, + loop=asyncio.get_running_loop(), + ) + try: + bus.send(can.Message()) + await asyncio.wait_for(event.wait(), 0.5) + self.assertEqual(len(collector.errors), 1) + self.assertIsInstance(collector.errors[0], ValueError) + self.assertIs(notifier.exception, collector.errors[0]) + finally: + notifier.stop() + + asyncio.run(run_it()) + + def test_async_listener_error_calls_on_error(self): + async def run_it(): + event = asyncio.Event() + collector = ErrorCollector(event) + + async def raising_callback(msg: can.Message) -> None: + raise RuntimeError("async listener failed") + + with can.Bus( + "async-listener-error", interface="virtual", receive_own_messages=True + ) as bus: + notifier = can.Notifier( + bus, + [raising_callback, collector], + 0.1, + loop=asyncio.get_running_loop(), + ) + try: + bus.send(can.Message()) + await asyncio.wait_for(event.wait(), 0.5) + self.assertEqual(len(collector.errors), 1) + self.assertIsInstance(collector.errors[0], RuntimeError) + self.assertIs(notifier.exception, collector.errors[0]) + finally: + notifier.stop() + + asyncio.run(run_it()) + if __name__ == "__main__": unittest.main()