diff --git a/src/apify/_actor.py b/src/apify/_actor.py index f846b0900..3841176e6 100644 --- a/src/apify/_actor.py +++ b/src/apify/_actor.py @@ -217,6 +217,9 @@ async def __aenter__(self) -> Self: # Initialize the event manager and register it in the service locator. await self.event_manager.__aenter__() + # Only the platform emits `ABORTING`, and it does so through `ApifyEventManager`. + if isinstance(self.event_manager, ApifyEventManager): + self.event_manager._on_internal(event=Event.ABORTING, listener=self._abort_child_runs) # noqa: SLF001 self.log.debug('Event manager initialized') # Initialize the charging manager. @@ -224,6 +227,7 @@ async def __aenter__(self) -> Self: await self._charging_manager_implementation.__aenter__() except BaseException: # Exit the already-entered event manager so its recurring tasks do not leak. + self._remove_internal_listeners() await self.event_manager.__aexit__(None, None, None) raise self.log.debug('Charging manager initialized') @@ -240,6 +244,7 @@ async def __aenter__(self) -> Self: except BaseException: # Undo the initialization, since a failed `__aenter__` gets no `__aexit__`. self._active = False + self._remove_internal_listeners() try: await self._charging_manager_implementation.__aexit__(None, None, None) finally: @@ -326,6 +331,7 @@ async def finalize() -> None: except TimeoutError: self.log.exception('Actor cleanup timed out') finally: + self._remove_internal_listeners() self._active = False if reraise_control_flow: @@ -988,6 +994,7 @@ async def start( force_permission_level: ActorPermissionLevel | None = None, webhooks: list[Webhook] | None = None, run_name: str | None = None, + abort_with_parent: bool = False, ) -> Run: """Run an Actor on the Apify platform. @@ -1023,10 +1030,16 @@ async def start( resurrected, and a new run is started only when nothing is recorded under the name, or the recorded run `FAILED` or no longer exists. The name is bound to the Actor and input it was first used with, so reusing it for a different Actor, task or input raises a `ValueError`. + abort_with_parent: If true, the child run is gracefully aborted when this Actor run is gracefully + aborted. It requires `run_name`, and the value is recorded under it, replacing the one from an earlier + call. A hard abort, a timeout or a crash of this Actor run leaves the child running. Returns: Info about the started Actor run """ + if abort_with_parent and run_name is None: + raise ValueError('`abort_with_parent` requires `run_name`, since only named child runs are tracked.') + if max_items is not None: _warn_max_items_deprecated() @@ -1062,6 +1075,7 @@ async def start( restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, timeout=timeout, + abort_with_parent=abort_with_parent, ) return run @@ -1176,6 +1190,7 @@ async def call( wait: timedelta | None = None, logger: logging.Logger | Literal['default'] | None = 'default', run_name: str | None = None, + abort_with_parent: bool = False, ) -> Run: """Start an Actor on the Apify Platform and wait for it to finish before returning. @@ -1214,10 +1229,16 @@ async def call( resurrected, and a new run is started only when nothing is recorded under the name, or the recorded run `FAILED` or no longer exists. The name is bound to the Actor and input it was first used with, so reusing it for a different Actor, task or input raises a `ValueError`. + abort_with_parent: If true, the child run is gracefully aborted when this Actor run is gracefully + aborted. It requires `run_name`, and the value is recorded under it, replacing the one from an earlier + call. A hard abort, a timeout or a crash of this Actor run leaves the child running. Returns: Info about the started Actor run. """ + if abort_with_parent and run_name is None: + raise ValueError('`abort_with_parent` requires `run_name`, since only named child runs are tracked.') + if max_items is not None: _warn_max_items_deprecated() @@ -1265,6 +1286,7 @@ async def call( restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, timeout=timeout, + abort_with_parent=abort_with_parent, ) # The earlier attempt of this call already streamed the log of a reattached or resurrected run. run = await self._wait_for_child_run( @@ -1291,6 +1313,7 @@ async def _find_or_start_child_run( restart_on_error: bool | None, memory_mbytes: int | None, timeout: timedelta | Literal['inherit'] | None, + abort_with_parent: bool, ) -> tuple[Run, bool]: return await self._child_run_registry.find_or_start( name, @@ -1307,8 +1330,16 @@ async def _find_or_start_child_run( max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, ), + abort_with_parent=abort_with_parent, ) + def _remove_internal_listeners(self) -> None: + if isinstance(self.event_manager, ApifyEventManager): + self.event_manager._off_internal(event=Event.ABORTING, listener=self._abort_child_runs) # noqa: SLF001 + + async def _abort_child_runs(self) -> None: + await self._child_run_registry.abort_runs_with_parent(self.apify_client) + async def _wait_for_child_run( self, name: str, @@ -1351,6 +1382,7 @@ async def start_task( webhooks: list[Webhook] | None = None, token: str | None = None, run_name: str | None = None, + abort_with_parent: bool = False, ) -> Run: """Start an Actor task on the Apify Platform. @@ -1386,10 +1418,16 @@ async def start_task( resurrected, and a new run is started only when nothing is recorded under the name, or the recorded run `FAILED` or no longer exists. The name is bound to the task and input it was first used with, so reusing it for a different Actor, task or input raises a `ValueError`. + abort_with_parent: If true, the child run is gracefully aborted when this Actor run is gracefully + aborted. It requires `run_name`, and the value is recorded under it, replacing the one from an earlier + call. A hard abort, a timeout or a crash of this Actor run leaves the child running. Returns: Info about the started Actor run. """ + if abort_with_parent and run_name is None: + raise ValueError('`abort_with_parent` requires `run_name`, since only named child runs are tracked.') + if max_items is not None: _warn_max_items_deprecated() @@ -1422,6 +1460,7 @@ async def start_task( restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, timeout=timeout, + abort_with_parent=abort_with_parent, ) return run @@ -1441,6 +1480,7 @@ async def call_task( wait: timedelta | None = None, token: str | None = None, run_name: str | None = None, + abort_with_parent: bool = False, ) -> Run: """Start an Actor task on the Apify Platform and wait for it to finish before returning. @@ -1476,10 +1516,16 @@ async def call_task( resurrected, and a new run is started only when nothing is recorded under the name, or the recorded run `FAILED` or no longer exists. The name is bound to the task and input it was first used with, so reusing it for a different Actor, task or input raises a `ValueError`. + abort_with_parent: If true, the child run is gracefully aborted when this Actor run is gracefully + aborted. It requires `run_name`, and the value is recorded under it, replacing the one from an earlier + call. A hard abort, a timeout or a crash of this Actor run leaves the child running. Returns: Info about the started Actor run. """ + if abort_with_parent and run_name is None: + raise ValueError('`abort_with_parent` requires `run_name`, since only named child runs are tracked.') + if max_items is not None: _warn_max_items_deprecated() @@ -1522,6 +1568,7 @@ async def call_task( restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, timeout=timeout, + abort_with_parent=abort_with_parent, ) run = await self._wait_for_child_run( run_name, client.run(started_run.id), started_run, wait=wait, logger=None, from_start=False diff --git a/src/apify/_child_runs.py b/src/apify/_child_runs.py index 8b5c7bc63..ec277dfd4 100644 --- a/src/apify/_child_runs.py +++ b/src/apify/_child_runs.py @@ -30,6 +30,8 @@ _RESURRECTABLE_STATUSES = frozenset({'ABORTED', 'TIMED-OUT'}) +_ABORTABLE_STATUSES = frozenset({'READY', 'RUNNING'}) + _NOT_FOUND_GRACE_SECS = 3 """How long a recorded run that the API reports as missing is looked up again before it counts as gone.""" @@ -60,6 +62,9 @@ class ChildRunRecord(ChildRunSnapshot): history: list[ChildRunSnapshot] = Field(default_factory=list) """Earlier runs under this name that failed or went missing and were replaced by a new run, oldest first.""" + abort_with_parent: bool = False + """Whether the current run is aborted when this Actor run is gracefully aborted.""" + def checksum_request(*, actor_id: str | None, task_id: str | None, run_input: Any) -> str: """Hash the Actor or task and the input of a named start, in the same JSON shape as the JS SDK.""" @@ -117,6 +122,7 @@ async def find_or_start( client: ApifyClientAsync, start_run: Callable[[], Awaitable[Run]], resurrect_run: Callable[[str], Awaitable[Run]], + abort_with_parent: bool = False, ) -> tuple[Run, bool]: """Return the run recorded under `name`, or start one when there is none to reuse. @@ -132,6 +138,8 @@ async def find_or_start( client: Client used to look up the recorded run. start_run: Starts a new run of the Actor or task. resurrect_run: Resurrects the recorded run, given its ID. + abort_with_parent: Whether to abort the run when this Actor run is gracefully aborted. It replaces the value + recorded under `name`. Returns: The run, and whether it was newly started. @@ -149,7 +157,9 @@ async def find_or_start( ) if record is None: - run = await self._start(name, checksum=checksum, start_run=start_run, history=[]) + run = await self._start( + name, checksum=checksum, start_run=start_run, history=[], abort_with_parent=abort_with_parent + ) self._clients[name] = client return run, True @@ -167,10 +177,17 @@ async def find_or_start( started_at=record.started_at, ) run = await self._start( - name, checksum=checksum, start_run=start_run, history=[*record.history, replaced] + name, + checksum=checksum, + start_run=start_run, + history=[*record.history, replaced], + abort_with_parent=abort_with_parent, ) return run, True + if record.abort_with_parent != abort_with_parent: + await self._save(name, record.model_copy(update={'abort_with_parent': abort_with_parent})) + if run.status in _RESURRECTABLE_STATUSES: logger.info(f'Resurrecting child run "{name}"', extra={'run_id': run.id, 'status': run.status}) run = await resurrect_run(run.id) @@ -207,6 +224,34 @@ async def update(self, name: str, run: Run) -> None: return await self._save(name, record.model_copy(update={'status': run.status})) + async def abort_runs_with_parent(self, client: ApifyClientAsync) -> None: + """Gracefully abort every recorded run marked `abort_with_parent` that is still `READY` or `RUNNING`. + + A failure to abort one run is logged and does not stop the others. + + Args: + client: Client used for a name not started or looked up in this process, e.g. after a migration. + """ + records = await self._load() + # Names with a start in flight are not recorded yet, so their locks are awaited too. + await asyncio.gather(*(self._abort(name, client) for name in {*records, *self._name_locks})) + + async def _abort(self, name: str, default_client: ApifyClientAsync) -> None: + async with self._name_locks.setdefault(name, asyncio.Lock()): + record = (await self._load()).get(name) + if record is None or not record.abort_with_parent: + return + run_client = self._clients.get(name, default_client).run(record.run_id) + try: + run = await run_client.get() + if run is None or run.status not in _ABORTABLE_STATUSES: + return + await run_client.abort(gracefully=True) + except Exception: + logger.exception(f'Failed to abort child run "{name}"', extra={'run_id': record.run_id}) + else: + logger.info(f'Aborted child run "{name}" with the parent', extra={'run_id': record.run_id}) + async def _start( self, name: str, @@ -214,10 +259,16 @@ async def _start( checksum: str, start_run: Callable[[], Awaitable[Run]], history: list[ChildRunSnapshot], + abort_with_parent: bool, ) -> Run: run = await start_run() record = ChildRunRecord( - run_id=run.id, status=run.status, started_at=run.started_at, checksum=checksum, history=history + run_id=run.id, + status=run.status, + started_at=run.started_at, + checksum=checksum, + history=history, + abort_with_parent=abort_with_parent, ) await self._save(name, record) return run diff --git a/src/apify/events/_apify_event_manager.py b/src/apify/events/_apify_event_manager.py index 4cb2524a1..6710b8c13 100644 --- a/src/apify/events/_apify_event_manager.py +++ b/src/apify/events/_apify_event_manager.py @@ -3,8 +3,9 @@ import asyncio import contextlib import time +from collections import defaultdict from logging import getLogger -from typing import TYPE_CHECKING, Annotated, Self, cast +from typing import TYPE_CHECKING, Annotated, Any, Self, cast import websockets.asyncio.client import websockets.client @@ -24,6 +25,7 @@ from types import TracebackType from crawlee.events._event_manager import EventManagerOptions + from crawlee.events._types import EventData, EventListener, WrappedListener from apify._configuration import Configuration @@ -94,6 +96,11 @@ def __init__(self, configuration: Configuration, **kwargs: Unpack[EventManagerOp connection, so that `__aenter__` can report it. """ + self._internal_listeners: defaultdict[Event, dict[EventListener[Any], WrappedListener]] = defaultdict(dict) + """Listeners of the SDK itself, mapped as `event -> listener -> wrapper`. `off` doesn't remove them, so user + code removing all listeners of an event keeps the SDK's own handling of it. + """ + @override async def __aenter__(self) -> Self: """Initialize the event manager upon entering the async context. @@ -149,6 +156,24 @@ async def __aexit__( # emitting `PersistState` again, as re-entering the context would be a no-op. await super().__aexit__(exc_type, exc_value, exc_traceback) + @override + def emit(self, *, event: Event, event_data: EventData) -> None: + super().emit(event=event, event_data=event_data) + + for listener, listener_wrapper in self._internal_listeners.get(event, {}).items(): + task_name = f'Task-{event.value}-{self._get_listener_name(listener)}' + listener_task = asyncio.create_task(listener_wrapper(event_data), name=task_name) + self._listener_tasks.add(listener_task) + listener_task.add_done_callback(self._listener_tasks.discard) + + def _on_internal(self, *, event: Event, listener: EventListener[Any]) -> None: + """Register a listener of the SDK itself, which `off` doesn't remove.""" + self._internal_listeners[event][listener] = self._wrap_listener(event, listener) + + def _off_internal(self, *, event: Event, listener: EventListener[Any]) -> None: + """Remove a listener registered by `_on_internal`.""" + self._internal_listeners.get(event, {}).pop(listener, None) + async def _teardown_platform_websocket(self) -> None: """Stop consuming the platform messages and close the websocket connection to the platform events.""" try: diff --git a/tests/e2e/test_actor_child_runs.py b/tests/e2e/test_actor_child_runs.py index f7de5d846..05b1352af 100644 --- a/tests/e2e/test_actor_child_runs.py +++ b/tests/e2e/test_actor_child_runs.py @@ -1,11 +1,15 @@ from __future__ import annotations import asyncio +from datetime import timedelta from typing import TYPE_CHECKING from apify import Actor +from apify._child_runs import CHILD_RUNS_KEY if TYPE_CHECKING: + from apify_client import ApifyClientAsync + from .conftest import MakeActorFunction, RunActorFunction @@ -88,3 +92,42 @@ async def main() -> None: assert run_result.status == 'SUCCEEDED' # The parent run and the one child run it resurrected. assert (await actor.runs().list()).total == 2 + + +async def test_named_child_run_is_aborted_with_parent( + make_actor: MakeActorFunction, + apify_client_async: ApifyClientAsync, +) -> None: + """A named child run started with `abort_with_parent` is aborted when the parent is gracefully aborted.""" + + async def main() -> None: + async with Actor: + actor_input = (await Actor.get_input()) or {} + if actor_input.get('is_child') is True: + await asyncio.sleep(300) + return + + actor_id = Actor.configuration.actor_id or '' + await Actor.start(actor_id=actor_id, run_input={'is_child': True}, run_name='child', abort_with_parent=True) + await asyncio.sleep(300) + + actor = await make_actor(label='child-run-abort-with-parent', main_func=main) + parent_run = await actor.start() + parent_kvs = apify_client_async.key_value_store(parent_run.default_key_value_store_id) + + # Wait for the parent to record the child run. + for _ in range(60): + if record := await parent_kvs.get_record(CHILD_RUNS_KEY): + break + await asyncio.sleep(2) + else: + raise AssertionError('The parent run did not record the child run in time.') + + child_run_id = record['value']['child']['runId'] + parent_run_client = apify_client_async.run(parent_run.id) + await parent_run_client.abort(gracefully=True) + await parent_run_client.wait_for_finish(wait_duration=timedelta(seconds=120)) + + child_run = await apify_client_async.run(child_run_id).wait_for_finish(wait_duration=timedelta(seconds=120)) + assert child_run is not None + assert child_run.status == 'ABORTED' diff --git a/tests/unit/actor/test_actor_child_runs.py b/tests/unit/actor/test_actor_child_runs.py index cce234666..043642b6d 100644 --- a/tests/unit/actor/test_actor_child_runs.py +++ b/tests/unit/actor/test_actor_child_runs.py @@ -3,15 +3,18 @@ import asyncio from datetime import timedelta from typing import TYPE_CHECKING, Any -from unittest.mock import MagicMock, Mock +from unittest.mock import AsyncMock, MagicMock, Mock import pytest from apify_client._models import Run +from crawlee import service_locator +from crawlee.events import Event, EventAbortingData -from apify import Actor, _child_runs +from apify import Actor, Configuration, _child_runs from apify._actor import _ActorType -from apify._child_runs import CHILD_RUNS_KEY, checksum_request +from apify._child_runs import CHILD_RUNS_KEY, ChildRunRegistry, checksum_request +from apify.events import ApifyEventManager if TYPE_CHECKING: from apify_client import ApifyClientAsync @@ -51,6 +54,7 @@ def stored_record( task_id: str | None = None, run_input: Any = None, history: list[dict[str, str]] | None = None, + abort_with_parent: bool = False, ) -> dict[str, Any]: """Build a registry record as it is stored in the default KVS.""" return { @@ -59,9 +63,18 @@ def stored_record( 'startedAt': STARTED_AT, 'checksum': checksum_request(actor_id=actor_id, task_id=task_id, run_input=run_input), 'history': history or [], + 'abortWithParent': abort_with_parent, } +@pytest.fixture +def apify_event_manager() -> ApifyEventManager: + """Make the Actor use `ApifyEventManager`, which delivers `ABORTING` on the platform, without a websocket.""" + event_manager = ApifyEventManager(Configuration.get_global_configuration()) + service_locator.set_event_manager(event_manager) + return event_manager + + async def record_child_run( name: str, run_id: str, @@ -722,3 +735,264 @@ async def test_child_runs_keeps_client_of_name_after_failed_lookup( default_http_client = Actor.apify_client._http_client assert child_runs['scrape-eu']._http_client is default_http_client + + +@pytest.mark.parametrize( + 'method', + [ + pytest.param('start', id='start'), + pytest.param('call', id='call'), + pytest.param('start_task', id='start task'), + pytest.param('call_task', id='call task'), + ], +) +async def test_abort_with_parent_requires_run_name( + apify_client_async_patcher: ApifyClientAsyncPatcher, method: str +) -> None: + """`abort_with_parent` without a `run_name` raises before any run is started.""" + for resource in ('actor', 'task'): + apify_client_async_patcher.patch(resource, 'start', return_value=make_run('new-run', 'READY')) + apify_client_async_patcher.patch(resource, 'call', return_value=make_run('new-run', 'SUCCEEDED')) + + async with Actor: + with pytest.raises(ValueError, match='requires `run_name`'): + await getattr(Actor, method)('some-id', abort_with_parent=True) + + for resource in ('actor', 'task'): + assert apify_client_async_patcher.calls[resource]['start'] == [] + assert apify_client_async_patcher.calls[resource]['call'] == [] + + +async def test_named_start_records_abort_with_parent(apify_client_async_patcher: ApifyClientAsyncPatcher) -> None: + """A named start with `abort_with_parent` records the flag with the run.""" + apify_client_async_patcher.patch('actor', 'start', return_value=make_run('new-run', 'READY')) + + async with Actor: + await Actor.start('some-actor', run_name='scrape-eu', abort_with_parent=True) + kvs = await Actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert stored['scrape-eu']['abortWithParent'] is True + + +async def test_named_call_task_records_abort_with_parent( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """A named task call with `abort_with_parent` records the flag with the run.""" + apify_client_async_patcher.patch('task', 'start', return_value=make_run('new-run', 'READY')) + apify_client_async_patcher.patch('run', 'wait_for_finish', return_value=make_run('new-run', 'SUCCEEDED')) + + async with Actor: + await Actor.call_task('some-task', run_name='scrape-eu', abort_with_parent=True) + kvs = await Actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert stored['scrape-eu']['abortWithParent'] is True + + +async def test_reattach_replaces_recorded_abort_with_parent( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """Reattaching under a name records the `abort_with_parent` value of the latest call.""" + apify_client_async_patcher.patch('run', 'get', return_value=make_run('old-run', 'RUNNING')) + + async with Actor: + await record_child_run('scrape-eu', 'old-run') + await Actor.start('some-actor', run_name='scrape-eu', abort_with_parent=True) + kvs = await Actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert stored['scrape-eu']['runId'] == 'old-run' + assert stored['scrape-eu']['abortWithParent'] is True + + +async def test_aborting_event_aborts_marked_active_child_runs( + apify_client_async_patcher: ApifyClientAsyncPatcher, apify_event_manager: ApifyEventManager +) -> None: + """On `ABORTING`, only child runs marked `abort_with_parent` that are still active are gracefully aborted.""" + runs = { + 'running-run': make_run('running-run', 'RUNNING'), + 'ready-run': make_run('ready-run', 'READY'), + 'finished-run': make_run('finished-run', 'SUCCEEDED'), + 'unmarked-run': make_run('unmarked-run', 'RUNNING'), + } + + async def get_run(run_client: Any, *_args: Any, **_kwargs: Any) -> Run | None: + return runs[run_client.resource_id] + + apify_client_async_patcher.patch('run', 'get', replacement_method=get_run) + apify_client_async_patcher.patch('run', 'abort', return_value=None) + + async with Actor: + kvs = await Actor.open_key_value_store() + await kvs.set_value( + CHILD_RUNS_KEY, + { + name: stored_record(run_id, 'RUNNING', abort_with_parent=marked) + for name, run_id, marked in [ + ('running', 'running-run', True), + ('ready', 'ready-run', True), + ('finished', 'finished-run', True), + ('unmarked', 'unmarked-run', False), + ] + }, + ) + await Actor._child_run_registry.load() + apify_event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await apify_event_manager.wait_for_all_listeners_to_complete() + + aborts = apify_client_async_patcher.calls['run']['abort'] + assert sorted(args[0].resource_id for args, _ in aborts) == ['ready-run', 'running-run'] + assert all(kwargs == {'gracefully': True} for _, kwargs in aborts) + + +async def test_failed_child_run_abort_does_not_stop_others( + apify_client_async_patcher: ApifyClientAsyncPatcher, + caplog: pytest.LogCaptureFixture, + apify_event_manager: ApifyEventManager, +) -> None: + """A child run that fails to abort is logged, and the other marked child runs are still aborted.""" + + async def abort_run(run_client: Any, *_args: Any, **_kwargs: Any) -> None: + if run_client.resource_id == 'broken-run': + raise RuntimeError('abort failed') + + apify_client_async_patcher.patch( + 'run', 'get', replacement_method=lambda run_client: make_run(run_client.resource_id, 'RUNNING') + ) + apify_client_async_patcher.patch('run', 'abort', replacement_method=abort_run) + + async with Actor: + kvs = await Actor.open_key_value_store() + await kvs.set_value( + CHILD_RUNS_KEY, + {name: stored_record(f'{name}-run', 'RUNNING', abort_with_parent=True) for name in ['broken', 'healthy']}, + ) + await Actor._child_run_registry.load() + apify_event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await apify_event_manager.wait_for_all_listeners_to_complete() + + aborts = apify_client_async_patcher.calls['run']['abort'] + assert sorted(args[0].resource_id for args, _ in aborts) == ['broken-run', 'healthy-run'] + assert 'Failed to abort child run "broken"' in caplog.text + assert 'Aborted child run "healthy" with the parent' in caplog.text + + +async def test_child_run_is_aborted_with_the_client_it_was_started_with() -> None: + """A child run started with its own client is aborted with that client, not the default one.""" + default_client = Mock() + child_client = Mock() + child_client.run.return_value.get = AsyncMock(return_value=make_run('new-run', 'RUNNING')) + child_client.run.return_value.abort = AsyncMock() + + async with Actor: + registry = ChildRunRegistry(Actor.open_key_value_store) + await registry.find_or_start( + 'scrape-eu', + actor_id='some-actor', + run_input=None, + client=child_client, + start_run=AsyncMock(return_value=make_run('new-run', 'READY')), + resurrect_run=AsyncMock(), + abort_with_parent=True, + ) + await registry.abort_runs_with_parent(default_client) + + child_client.run.return_value.abort.assert_awaited_once_with(gracefully=True) + default_client.run.assert_not_called() + + +async def test_rejected_named_start_keeps_the_client_used_to_abort() -> None: + """A named start rejected for another Actor does not change the client its recorded run is aborted with.""" + default_client = Mock() + default_client.run.return_value.get = AsyncMock(return_value=make_run('old-run', 'RUNNING')) + default_client.run.return_value.abort = AsyncMock() + other_client = Mock() + + async with Actor: + kvs = await Actor.open_key_value_store() + await kvs.set_value(CHILD_RUNS_KEY, {'scrape-eu': stored_record('old-run', 'RUNNING', abort_with_parent=True)}) + registry = ChildRunRegistry(Actor.open_key_value_store) + with pytest.raises(ValueError, match='already used for a different Actor, task or input'): + await registry.find_or_start( + 'scrape-eu', + actor_id='other-actor', + run_input=None, + client=other_client, + start_run=AsyncMock(), + resurrect_run=AsyncMock(), + ) + await registry.abort_runs_with_parent(default_client) + + default_client.run.return_value.abort.assert_awaited_once_with(gracefully=True) + other_client.run.assert_not_called() + + +async def test_aborting_waits_for_a_named_start_in_flight() -> None: + """A named start in flight when the parent is aborted has its run aborted once the run is recorded.""" + client = Mock() + client.run.return_value.get = AsyncMock(return_value=make_run('new-run', 'RUNNING')) + client.run.return_value.abort = AsyncMock() + started = asyncio.Event() + release = asyncio.Event() + + async def start_run() -> Run: + started.set() + await release.wait() + return make_run('new-run', 'READY') + + async with Actor: + registry = ChildRunRegistry(Actor.open_key_value_store) + start_task = asyncio.create_task( + registry.find_or_start( + 'scrape-eu', + actor_id='some-actor', + run_input=None, + client=client, + start_run=start_run, + resurrect_run=AsyncMock(), + abort_with_parent=True, + ) + ) + await started.wait() + abort_task = asyncio.create_task(registry.abort_runs_with_parent(client)) + await asyncio.sleep(0) + assert not abort_task.done() + release.set() + await asyncio.gather(start_task, abort_task) + + client.run.return_value.abort.assert_awaited_once_with(gracefully=True) + + +async def test_exit_removes_the_aborting_listener( + apify_client_async_patcher: ApifyClientAsyncPatcher, apify_event_manager: ApifyEventManager +) -> None: + """After the Actor exits, an `ABORTING` event on a still-active event manager aborts no child run.""" + apify_client_async_patcher.patch('actor', 'start', return_value=make_run('new-run', 'READY')) + apify_client_async_patcher.patch('run', 'get', return_value=make_run('new-run', 'RUNNING')) + apify_client_async_patcher.patch('run', 'abort', return_value=None) + + async with apify_event_manager: + async with Actor: + await Actor.start('some-actor', run_name='scrape-eu', abort_with_parent=True) + apify_event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await apify_event_manager.wait_for_all_listeners_to_complete() + + assert apify_client_async_patcher.calls['run']['abort'] == [] + + +async def test_removing_all_aborting_listeners_keeps_aborting_child_runs( + apify_client_async_patcher: ApifyClientAsyncPatcher, apify_event_manager: ApifyEventManager +) -> None: + """Removing all `ABORTING` listeners from the event manager still aborts child runs marked `abort_with_parent`.""" + apify_client_async_patcher.patch('actor', 'start', return_value=make_run('new-run', 'READY')) + apify_client_async_patcher.patch('run', 'get', return_value=make_run('new-run', 'RUNNING')) + apify_client_async_patcher.patch('run', 'abort', return_value=None) + + async with Actor: + await Actor.start('some-actor', run_name='scrape-eu', abort_with_parent=True) + apify_event_manager.off(event=Event.ABORTING) + apify_event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await apify_event_manager.wait_for_all_listeners_to_complete() + + assert len(apify_client_async_patcher.calls['run']['abort']) == 1 diff --git a/tests/unit/events/test_apify_event_manager.py b/tests/unit/events/test_apify_event_manager.py index 6d6ec65c0..11f13d0f4 100644 --- a/tests/unit/events/test_apify_event_manager.py +++ b/tests/unit/events/test_apify_event_manager.py @@ -9,14 +9,14 @@ from collections import defaultdict from datetime import timedelta from typing import TYPE_CHECKING, Any -from unittest.mock import Mock +from unittest.mock import AsyncMock, Mock import pytest import websockets import websockets.asyncio.client import websockets.asyncio.server -from crawlee.events._types import Event +from crawlee.events._types import Event, EventAbortingData from ..._utils import poll_until_condition from apify import Configuration @@ -489,6 +489,32 @@ async def handler(_data: Any) -> None: assert persist_state_counter == 0 +async def test_internal_listener_is_kept_by_off() -> None: + """A listener registered by `_on_internal` still runs after `off` removes all listeners of its event.""" + listener = AsyncMock() + + async with ApifyEventManager(Configuration.get_global_configuration()) as event_manager: + event_manager._on_internal(event=Event.ABORTING, listener=listener) + event_manager.off(event=Event.ABORTING) + event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await event_manager.wait_for_all_listeners_to_complete() + + listener.assert_awaited_once() + + +async def test_off_internal_removes_the_listener() -> None: + """A listener removed by `_off_internal` no longer runs on its event.""" + listener = AsyncMock() + + async with ApifyEventManager(Configuration.get_global_configuration()) as event_manager: + event_manager._on_internal(event=Event.ABORTING, listener=listener) + event_manager._off_internal(event=Event.ABORTING, listener=listener) + event_manager.emit(event=Event.ABORTING, event_data=EventAbortingData()) + await event_manager.wait_for_all_listeners_to_complete() + + listener.assert_not_awaited() + + async def test_deprecated_event_is_skipped(monkeypatch: pytest.MonkeyPatch) -> None: """Test that deprecated events (like CPU_INFO) are silently skipped.""" async with (