Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 27 additions & 2 deletions can/notifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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.
Expand Down
70 changes: 70 additions & 0 deletions test/notifier_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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()