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
38 changes: 26 additions & 12 deletions src/crawlee/events/_event_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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.
"""

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)
for result in results:
if isinstance(result, Exception):
logger.error('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)
# 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. Only a waiter
# that is a listener itself can be awaited this way, so only such waiters are tracked and excluded.
waiting_task = asyncio.current_task()
is_listener_waiter = waiting_task is not None and waiting_task in self._listener_tasks
if is_listener_waiter and waiting_task is not None:
self._waiting_listener_tasks.add(waiting_task)

# `emit` only schedules the listener wrappers; each registers its listener task once it starts running,
# so yield first - otherwise listeners of a just-emitted event are missed and the wait is a no-op.
await asyncio.sleep(0)

# Any other caller is outside the cycle, so it must await every listener, waiting ones included.
excluded = self._waiting_listener_tasks if is_listener_waiter else frozenset[asyncio.Task]()
listener_tasks = [task for task in self._listener_tasks if task not in excluded]

try:
await wait_for_all_tasks_for_finish(tasks=listener_tasks, logger=logger, timeout=timeout)
finally:
if is_listener_waiter and 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."""
Expand Down
141 changes: 141 additions & 0 deletions tests/unit/events/test_event_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -171,6 +172,7 @@ async def test_close_after_emit_processes_event(
async def test_wait_for_all_listeners_cancelled_error(
monkeypatch: pytest.MonkeyPatch,
caplog: pytest.LogCaptureFixture,
event_system_info_data: EventSystemInfoData,
) -> None:
# Simulate long-running listener tasks
async def long_running_listener() -> None:
Expand All @@ -183,6 +185,7 @@ async def mock_async_wait(*_: Any, **__: Any) -> None:
with pytest.raises(asyncio.CancelledError), caplog.at_level(logging.WARNING): # noqa: PT012
async with EventManager(close_timeout=timedelta(milliseconds=10)) as event_manager:
event_manager.on(event=Event.SYSTEM_INFO, listener=long_running_listener)
event_manager.emit(event=Event.SYSTEM_INFO, event_data=event_system_info_data)

# Use monkeypatch to replace asyncio.wait with mock_async_wait
monkeypatch.setattr('asyncio.wait', mock_async_wait)
Expand All @@ -206,6 +209,144 @@ 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_wait_from_outside_awaits_a_listener_that_is_itself_waiting(
event_manager: EventManager,
event_system_info_data: EventSystemInfoData,
) -> None:
"""A caller that is not a listener is outside the deadlock cycle, so it must await even waiting listeners."""
parked = asyncio.Event()
waiting_listener_done = asyncio.Event()

async def other_listener(_: Any) -> None:
await asyncio.sleep(0.1)

async def waiting_listener(_: Any) -> None:
parked.set()
await event_manager.wait_for_all_listeners_to_complete()
# Work done after the inner wait returns - the outer wait must not return before it finishes.
await asyncio.sleep(0.1)
waiting_listener_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)

# `parked` is set right before the listener registers itself as a waiter, without yielding in between.
await asyncio.wait_for(parked.wait(), timeout=5)

await asyncio.wait_for(event_manager.wait_for_all_listeners_to_complete(), timeout=5)

assert waiting_listener_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()
Expand Down
Loading