diff --git a/src/crawlee/events/_event_manager.py b/src/crawlee/events/_event_manager.py index d44dd0f8c4..2132426198 100644 --- a/src/crawlee/events/_event_manager.py +++ b/src/crawlee/events/_event_manager.py @@ -81,6 +81,9 @@ def __init__( # Listeners are wrapped inside asyncio.Task. Store their references here so that we can wait for them to finish. self._listener_tasks: set[asyncio.Task] = set() + # Tasks currently blocked in `wait_for_all_listeners_to_complete`; excluded when gathering to avoid deadlock. + self._waiting_listener_tasks: set[asyncio.Task] = set() + # Store the mapping between events, listeners and their wrappers in the following way: # event -> listener -> [wrapped_listener_1, wrapped_listener_2, ...] self._listeners_to_wrappers: dict[Event, dict[EventListener[Any], list[WrappedListener]]] = defaultdict( @@ -202,7 +205,8 @@ async def listener_wrapper(event_data: EventData) -> None: ) finally: logger.debug('EventManager.on.listener_wrapper(): Removing listener task from the set...') - self._listener_tasks.remove(listener_task) + # `discard`, not `remove`: `__aexit__` may have cleared the set while this listener ran. + self._listener_tasks.discard(listener_task) self._listeners_to_wrappers[event][listener].append(listener_wrapper) self._event_emitter.add_listener(event.value, listener_wrapper) @@ -256,17 +260,27 @@ async def wait_for_all_listeners_to_complete(self, *, timeout: timedelta | None timeout: The maximum time to wait for the event listeners to finish. If they do not complete within the specified timeout, they will be canceled. """ + # A waiter can't finish until the listeners it awaits do, so waiters must never await each other or + # themselves - this is what happens when a listener waits or closes from within itself. + waiting_task = asyncio.current_task() + if waiting_task is not None: + self._waiting_listener_tasks.add(waiting_task) async def wait_for_listeners() -> None: """Gathers all listener tasks and awaits their completion, logging any exceptions encountered.""" - results = await asyncio.gather(*self._listener_tasks, return_exceptions=True) + listener_tasks = [task for task in self._listener_tasks if task not in self._waiting_listener_tasks] + results = await asyncio.gather(*listener_tasks, return_exceptions=True) for result in results: if isinstance(result, Exception): logger.exception('Event listener raised an exception.', exc_info=result) tasks = [asyncio.create_task(wait_for_listeners(), name=f'Task-{wait_for_listeners.__name__}')] - await wait_for_all_tasks_for_finish(tasks=tasks, logger=logger, timeout=timeout) + try: + await wait_for_all_tasks_for_finish(tasks=tasks, logger=logger, timeout=timeout) + finally: + if waiting_task is not None: + self._waiting_listener_tasks.discard(waiting_task) async def _emit_persist_state_event(self) -> None: """Emit a persist state event with the given migration status.""" diff --git a/tests/unit/events/test_event_manager.py b/tests/unit/events/test_event_manager.py index 2b7b7cd10c..584bbe9ced 100644 --- a/tests/unit/events/test_event_manager.py +++ b/tests/unit/events/test_event_manager.py @@ -2,6 +2,7 @@ import asyncio import logging +from contextlib import suppress from datetime import timedelta from functools import update_wrapper from typing import TYPE_CHECKING, Any @@ -206,6 +207,114 @@ async def test_methods_raise_error_when_not_active(event_system_info_data: Event assert event_manager.active is True +async def test_wait_for_all_listeners_from_within_a_listener_does_not_deadlock( + event_manager: EventManager, + event_system_info_data: EventSystemInfoData, +) -> None: + """Waiting from within a listener must not self-await, yet must still await the other listeners.""" + other_listener_done = asyncio.Event() + waiter_done = asyncio.Event() + other_done_when_wait_returned: bool | None = None + + async def other_listener(_: Any) -> None: + await asyncio.sleep(0.2) + other_listener_done.set() + + async def waiting_listener(_: Any) -> None: + nonlocal other_done_when_wait_returned + await event_manager.wait_for_all_listeners_to_complete() + other_done_when_wait_returned = other_listener_done.is_set() + waiter_done.set() + + event_manager.on(event=Event.SYSTEM_INFO, listener=other_listener) + event_manager.on(event=Event.SYSTEM_INFO, listener=waiting_listener) + event_manager.emit(event=Event.SYSTEM_INFO, event_data=event_system_info_data) + + await asyncio.wait_for(waiter_done.wait(), timeout=5) + + # No self-await deadlock, and the wait must have blocked until the co-registered listener finished. + assert other_done_when_wait_returned is True + assert other_listener_done.is_set() + + +async def test_wait_from_within_multiple_listeners_does_not_deadlock( + event_manager: EventManager, + event_system_info_data: EventSystemInfoData, +) -> None: + """Several listeners each waiting for all listeners at once must not deadlock one another.""" + first_done = asyncio.Event() + second_done = asyncio.Event() + + async def first_waiting_listener(_: Any) -> None: + await event_manager.wait_for_all_listeners_to_complete() + first_done.set() + + async def second_waiting_listener(_: Any) -> None: + await event_manager.wait_for_all_listeners_to_complete() + second_done.set() + + event_manager.on(event=Event.SYSTEM_INFO, listener=first_waiting_listener) + event_manager.on(event=Event.SYSTEM_INFO, listener=second_waiting_listener) + event_manager.emit(event=Event.SYSTEM_INFO, event_data=event_system_info_data) + + await asyncio.wait_for(asyncio.gather(first_done.wait(), second_done.wait()), timeout=5) + + assert first_done.is_set() + assert second_done.is_set() + + +async def test_close_from_within_a_listener_does_not_deadlock_or_error( + event_system_info_data: EventSystemInfoData, +) -> None: + """Closing the event manager from within a listener (as `Actor.exit()` does) must not deadlock or raise.""" + event_manager = EventManager() + await event_manager.__aenter__() + + # A wrapper finalizing after close raises onto the loop (its `error` listener is gone by then), so watch + # both channels for a stray exception. + emitter_errors: list[BaseException] = [] + event_manager._event_emitter.add_listener('error', emitter_errors.append) + loop_errors: list[dict[str, Any]] = [] + asyncio.get_running_loop().set_exception_handler(lambda _loop, context: loop_errors.append(context)) + + closed = asyncio.Event() + other_listener_done = asyncio.Event() + + async def other_listener(_: Any) -> None: + await asyncio.sleep(0.2) + other_listener_done.set() + + async def closing_listener(_: Any) -> None: + await event_manager.__aexit__(None, None, None) + closed.set() + + # A second listener makes close await a concurrently-running listener - the real `Actor.exit()` shape. + event_manager.on(event=Event.SYSTEM_INFO, listener=other_listener) + event_manager.on(event=Event.SYSTEM_INFO, listener=closing_listener) + + tasks_before = asyncio.all_tasks() + event_manager.emit(event=Event.SYSTEM_INFO, event_data=event_system_info_data) + + try: + await asyncio.wait_for(closed.wait(), timeout=5) + # Drain the wrapper tasks so their `finally` blocks run before we assert - no arbitrary sleep. + spawned = asyncio.all_tasks() - tasks_before - {asyncio.current_task()} + if spawned: + await asyncio.wait(spawned) + finally: + # Cap the cleanup so a regressed deadlock surfaces the real failure instead of hanging. + if event_manager.active: + with suppress(Exception): + await asyncio.wait_for(event_manager.__aexit__(None, None, None), timeout=5) + + # With `discard` no wrapper raises on finalize; the `remove` regression would surface on one of these. + assert emitter_errors == [] + assert loop_errors == [] + assert other_listener_done.is_set() + assert event_manager.active is False + assert len(event_manager._listener_tasks) == 0 + + async def test_event_manager_in_context_persistence() -> None: """Test that entering the `EventManager` context emits persist state event at least once.""" event_manager = EventManager()