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
47 changes: 47 additions & 0 deletions src/apify/_actor.py
Original file line number Diff line number Diff line change
Expand Up @@ -217,13 +217,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 All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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()

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

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

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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(
Expand All @@ -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,
Expand All @@ -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,
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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()

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

Expand All @@ -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.

Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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
Expand Down
57 changes: 54 additions & 3 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'})

_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."""

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

Expand All @@ -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.
Expand All @@ -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

Expand All @@ -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)
Expand Down Expand Up @@ -207,17 +224,51 @@ 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,
*,
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
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