Skip to content
Draft
38 changes: 38 additions & 0 deletions src/apify/_actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,13 +206,17 @@ 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.
try:
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')
Expand Down Expand Up @@ -304,6 +308,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:
Expand Down Expand Up @@ -947,6 +952,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.

Expand Down Expand Up @@ -979,10 +985,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 `actor_id` exactly as passed, so reusing it with
any other value, or for a task, 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.')

client = self.new_client(token=token) if token else self.apify_client

if timeout == 'inherit':
Expand Down Expand Up @@ -1021,6 +1033,7 @@ async def start(
restart_on_error=restart_on_error,
memory_mbytes=memory_mbytes,
run_timeout=actor_start_timeout,
abort_with_parent=abort_with_parent,
)
return run

Expand Down Expand Up @@ -1079,6 +1092,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.

Expand Down Expand Up @@ -1114,10 +1128,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 `actor_id` exactly as passed, so reusing it with
any other value, or for a task, 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.')

client = self.new_client(token=token) if token else self.apify_client

if timeout == 'inherit':
Expand Down Expand Up @@ -1167,6 +1187,7 @@ async def call(
restart_on_error=restart_on_error,
memory_mbytes=memory_mbytes,
run_timeout=actor_call_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(
Expand All @@ -1191,6 +1212,7 @@ async def _find_or_start_child_run(
restart_on_error: bool | None,
memory_mbytes: int | None,
run_timeout: timedelta | None,
abort_with_parent: bool,
) -> tuple[Run, bool]:
return await self._child_run_registry.find_or_start(
name,
Expand All @@ -1205,8 +1227,16 @@ async def _find_or_start_child_run(
memory_mbytes=memory_mbytes,
run_timeout=run_timeout,
),
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,
run_client: RunClientAsync,
Expand Down Expand Up @@ -1256,6 +1286,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.

Expand Down Expand Up @@ -1287,10 +1318,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 `task_id` exactly as passed, so reusing it with
any other value, or for an Actor, 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.')

client = self.new_client(token=token) if token else self.apify_client

if timeout == 'inherit':
Expand Down Expand Up @@ -1333,6 +1370,7 @@ async def call_task(
restart_on_error=restart_on_error,
memory_mbytes=memory_mbytes,
run_timeout=task_call_timeout,
abort_with_parent=abort_with_parent,
)
run = await self._wait_for_child_run(
client.run(started_run.id), started_run, wait=wait, logger=None, from_start=False
Expand Down
85 changes: 74 additions & 11 deletions src/apify/_child_runs.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,8 @@

_RESURRECTABLE_STATUSES = frozenset({'ABORTED', 'TIMED-OUT'})

_ABORTABLE_STATUSES = frozenset({'READY', 'RUNNING'})


class ChildRunRecord(BaseModel):
"""A child run tracked under a name in the child run registry."""
Expand All @@ -48,6 +50,9 @@ class ChildRunRecord(BaseModel):
previous_run_ids: list[str] = Field(default_factory=list)
"""IDs of 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."""

@model_validator(mode='after')
def _check_started_from(self) -> Self:
if (self.actor_id is None) == (self.task_id is None):
Expand Down Expand Up @@ -79,6 +84,9 @@ class ChildRunInfo:
previous_run_ids: list[str]
"""IDs of earlier runs under this name that failed or went missing and were replaced by a new run, oldest first."""

abort_with_parent: bool
"""Whether the current run is aborted when this Actor run is gracefully aborted."""


_records_adapter = TypeAdapter(dict[str, ChildRunRecord])

Expand All @@ -96,6 +104,8 @@ def __init__(self, open_key_value_store: Callable[[], Awaitable[KeyValueStore]])
self._load_lock = asyncio.Lock()
self._write_lock = asyncio.Lock()
self._name_locks: defaultdict[str, asyncio.Lock] = defaultdict(asyncio.Lock)
self._clients: dict[str, ApifyClientAsync] = {}
"""Client each name was last started or reattached with in this process, used to abort its run."""

async def find_or_start(
self,
Expand All @@ -106,6 +116,7 @@ async def find_or_start(
client: ApifyClientAsync,
start_run: Callable[[], Awaitable[Run]],
resurrect_run: Callable[[RunClientAsync], 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.

Expand All @@ -117,9 +128,11 @@ async def find_or_start(
name: Name of the child run, unique within the parent run.
actor_id: The Actor to start. It must match the Actor already recorded under `name`.
task_id: The task to start, in place of `actor_id`. It must match the task already recorded under `name`.
client: Client used to look up and resurrect the recorded run.
client: Client used to look up, resurrect and abort the recorded run.
start_run: Starts a new run of the Actor or task.
resurrect_run: Resurrects the recorded run, given its run client.
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.
Expand All @@ -128,32 +141,46 @@ async def find_or_start(
records = await self._load()
record = records.get(name)

if record is None:
run = await self._start(
name, actor_id=actor_id, task_id=task_id, start_run=start_run, previous_run_ids=[]
)
return run, True

if (record.actor_id, record.task_id) != (actor_id, task_id):
if record is not None and (record.actor_id, record.task_id) != (actor_id, task_id):
raise ValueError(
f'Child run "{name}" is already recorded for '
f'{_describe_started_from(record.actor_id, record.task_id)}, '
f'it cannot be reused for {_describe_started_from(actor_id, task_id)}.'
)

self._clients[name] = client

if record is None:
run = await self._start(
name,
actor_id=actor_id,
task_id=task_id,
start_run=start_run,
previous_run_ids=[],
abort_with_parent=abort_with_parent,
)
return run, True

run_client = client.run(record.run_id)
run = await run_client.get()

if run is not None and run.status in _SETTLING_STATUSES:
run = await run_client.wait_for_finish()

if run is None or run.status == 'FAILED':
previous_run_ids = [*record.previous_run_ids, record.run_id]
run = await self._start(
name, actor_id=actor_id, task_id=task_id, start_run=start_run, previous_run_ids=previous_run_ids
name,
actor_id=actor_id,
task_id=task_id,
start_run=start_run,
previous_run_ids=[*record.previous_run_ids, record.run_id],
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})
return await resurrect_run(run_client), False
Expand All @@ -177,10 +204,39 @@ async def list_runs(self, client: ApifyClientAsync) -> dict[str, ChildRunInfo]:
run_id=record.run_id,
run=run,
previous_run_ids=list(record.previous_run_ids),
abort_with_parent=record.abort_with_parent,
)
for (name, record), run in zip(records.items(), runs, strict=True)
}

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 reattached 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[name]:
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,
Expand All @@ -189,9 +245,16 @@ async def _start(
task_id: str | None,
start_run: Callable[[], Awaitable[Run]],
previous_run_ids: list[str],
abort_with_parent: bool,
) -> Run:
run = await start_run()
record = ChildRunRecord(actor_id=actor_id, task_id=task_id, run_id=run.id, previous_run_ids=previous_run_ids)
record = ChildRunRecord(
actor_id=actor_id,
task_id=task_id,
run_id=run.id,
previous_run_ids=previous_run_ids,
abort_with_parent=abort_with_parent,
)
await self._save(name, record)
return run

Expand Down
27 changes: 26 additions & 1 deletion src/apify/events/_apify_event_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down
Loading
Loading