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
20 changes: 17 additions & 3 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.
"""
# 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."""
Expand Down
109 changes: 109 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 @@ -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()
Expand Down
Loading