diff --git a/src/apify/_actor.py b/src/apify/_actor.py index b662c5aa..a01845a4 100644 --- a/src/apify/_actor.py +++ b/src/apify/_actor.py @@ -35,7 +35,7 @@ ChargingManagerImplementation, charge_lock_if_charging, ) -from apify._child_runs import ChildRunInfo, ChildRunRegistry +from apify._child_runs import ChildRunInfo, ChildRunRegistry, StartRun from apify._configuration import Configuration from apify._consts import EVENT_LISTENERS_TIMEOUT, EXIT_CODE_ERROR_USER_FUNCTION_THREW, ActorEnvVars, ApifyEnvVars from apify._crypto import decrypt_input_secrets, load_private_key @@ -50,7 +50,7 @@ if TYPE_CHECKING: import logging - from collections.abc import Awaitable, Callable, MutableMapping + from collections.abc import Callable, MutableMapping from decimal import Decimal from types import TracebackType from typing import Self @@ -153,7 +153,9 @@ def __init__( # Keep track of all used state stores to persist their values on exit self._use_state_stores: set[str | None] = set() - self._child_run_registry = ChildRunRegistry(self.open_key_value_store) + self._child_run_registry = ChildRunRegistry( + self.open_key_value_store, lambda: self._charging_manager_implementation + ) self._active = False """Whether the Actor instance is currently active (initialized and within context).""" @@ -212,6 +214,7 @@ async def __aenter__(self) -> Self: self.log.debug('Event manager initialized') # Initialize the charging manager. + self._charging_manager_implementation.child_run_reservations = self._child_run_registry.reserved_usd try: await self._charging_manager_implementation.__aenter__() except BaseException: @@ -224,6 +227,10 @@ async def __aenter__(self) -> Self: # Mark initialization as complete and update global state. self._active = True + # Child runs recorded by an earlier attempt of this run keep their part of the budget reserved. + if self._charging_manager_implementation.get_max_total_charge_usd().is_finite(): + await self._child_run_registry.load() + if not Actor.is_at_home(): # Make sure that the input related KVS is initialized to ensure that the input aware client is used await self.open_key_value_store() @@ -966,7 +973,11 @@ async def start( content_type: The content type of the input. build: Specifies the Actor build to run. It can be either a build tag or build number. By default, the run uses the build specified in the default run configuration for the Actor (typically latest). - max_total_charge_usd: A limit on the total charged amount for pay-per-event Actors. + max_total_charge_usd: A limit on the total charged amount for pay-per-event Actors. When `run_name` is + set and this Actor run was started with a `max_total_charge_usd` set by the user, the limit defaults to + the part of that budget not charged by this Actor run nor reserved for its other named child runs, + and a higher value is lowered to it. The limit stays reserved until the child run finishes and its + charge is known. restart_on_error: If true, the Actor run process will be restarted whenever it exits with a non-zero status code. memory_mbytes: Memory limit for the run, in megabytes. By default, the run uses a memory limit specified @@ -1012,7 +1023,6 @@ async def start( run_input=run_input, content_type=content_type, build=build, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=actor_start_timeout, @@ -1021,7 +1031,7 @@ async def start( ) if run_name is None: - return await start_run() + return await start_run(max_total_charge_usd=max_total_charge_usd) run, _ = await self._find_or_start_child_run( run_name, @@ -1105,7 +1115,11 @@ async def call( content_type: The content type of the input. build: Specifies the Actor build to run. It can be either a build tag or build number. By default, the run uses the build specified in the default run configuration for the Actor (typically latest). - max_total_charge_usd: A limit on the total charged amount for pay-per-event Actors. + max_total_charge_usd: A limit on the total charged amount for pay-per-event Actors. When `run_name` is + set and this Actor run was started with a `max_total_charge_usd` set by the user, the limit defaults to + the part of that budget not charged by this Actor run nor reserved for its other named child runs, + and a higher value is lowered to it. The limit stays reserved until the child run finishes and its + charge is known. restart_on_error: If true, the Actor run process will be restarted whenever it exits with a non-zero status code. memory_mbytes: Memory limit for the run, in megabytes. By default, the run uses a memory limit specified @@ -1175,7 +1189,6 @@ async def call( run_input=run_input, content_type=content_type, build=build, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=actor_call_timeout, @@ -1208,7 +1221,7 @@ async def _find_or_start_child_run( actor_id: str | None = None, task_id: str | None = None, client: ApifyClientAsync, - start_run: Callable[[], Awaitable[Run]], + start_run: StartRun, build: str | None, max_total_charge_usd: Decimal | None, restart_on_error: bool | None, @@ -1222,7 +1235,7 @@ async def _find_or_start_child_run( task_id=task_id, client=client, start_run=start_run, - resurrect_run=lambda run_client: run_client.resurrect( + resurrect_run=lambda run_client, max_total_charge_usd: run_client.resurrect( build=build, max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, @@ -1230,6 +1243,7 @@ async def _find_or_start_child_run( run_timeout=run_timeout, ), abort_with_parent=abort_with_parent, + max_total_charge_usd=max_total_charge_usd, ) def _remove_internal_listeners(self) -> None: @@ -1335,7 +1349,9 @@ async def call_task( to the recorded run. A `SUCCEEDED` run is returned as is, an `ABORTED` or `TIMED-OUT` one is 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`. + any other value, or for an Actor, raises a `ValueError`. When this Actor run was started with + a `max_total_charge_usd` set by the user, a named call that would start a new run raises + a `RuntimeError`, since the task's run cannot be given a part of that budget. 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. @@ -1370,19 +1386,27 @@ async def call_task( wait_duration=wait, ) else: - started_run, _ = await self._find_or_start_child_run( - run_name, - task_id=task_id, - client=client, - start_run=partial( - task_client.start, + + async def start_task_run(*, max_total_charge_usd: Decimal | None) -> Run: + if max_total_charge_usd is not None: + raise RuntimeError( + f'Child run "{run_name}" was not started, since a task run cannot be given a part of the ' + 'budget of this Actor run.' + ) + return await task_client.start( task_input=task_input, build=build, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=task_call_timeout, webhooks=to_client_representations(webhooks), - ), + ) + + started_run, _ = await self._find_or_start_child_run( + run_name, + task_id=task_id, + client=client, + start_run=start_task_run, build=build, max_total_charge_usd=None, restart_on_error=restart_on_error, diff --git a/src/apify/_charging.py b/src/apify/_charging.py index aebc1c1f..8bb7cd94 100644 --- a/src/apify/_charging.py +++ b/src/apify/_charging.py @@ -25,7 +25,7 @@ from apify.storages import Dataset if TYPE_CHECKING: - from collections.abc import AsyncIterator + from collections.abc import AsyncIterator, Callable from types import TracebackType from apify_client import ApifyClientAsync @@ -343,6 +343,10 @@ def __init__(self, configuration: Configuration, client: ApifyClientAsync) -> No self.charge_lock = ReentrantLock() + self.child_run_reservations: Callable[[], Decimal] = Decimal + """Returns the part of `max_total_charge_usd` reserved for child runs of this Actor run.""" + self._is_max_total_charge_usd_set_by_user: bool | None = None + async def __aenter__(self) -> None: """Initialize the charging manager - this is called by the `Actor` class and shouldn't be invoked manually.""" # Validate config @@ -563,9 +567,33 @@ def calculate_max_event_charge_count_within_limit(self, event_name: str) -> int if not price: return None - result = (self._max_total_charge_usd - self.calculate_total_charged_amount()) / price + result = self.calculate_remaining_budget() / price return max(0, math.floor(result)) if result.is_finite() else None + @_ensure_context + def calculate_remaining_budget(self) -> Decimal: + """Return the part of `max_total_charge_usd` not charged by this Actor run nor reserved for its child runs.""" + return self._max_total_charge_usd - self.calculate_total_charged_amount() - self.child_run_reservations() + + @_ensure_context + async def is_max_total_charge_usd_set_by_user(self) -> bool: + """Return whether `max_total_charge_usd` was set for this Actor run, not defaulted by the platform. + + The platform gives pay-per-event runs a limit even when nobody set one, and marks the run options when the + limit was set. A run that does not say so is treated as having a default limit. + """ + if not self._max_total_charge_usd.is_finite(): + return False + if not self._is_at_home: + return True + if self._is_max_total_charge_usd_set_by_user is None: + if self._actor_run_id is None: + raise RuntimeError('Actor run ID not configured') + run = await self._client.run(self._actor_run_id).get() + extra = (run.options.model_extra or {}) if run is not None else {} + self._is_max_total_charge_usd_set_by_user = extra.get('isMaxTotalChargeUsdSetByUser') is True + return self._is_max_total_charge_usd_set_by_user + @_ensure_context def get_pricing_info(self) -> ActorPricingInfo: return ActorPricingInfo( @@ -603,7 +631,7 @@ def compute_push_data_limit( if not combined_price: return items_count - result = (self._max_total_charge_usd - self.calculate_total_charged_amount()) / combined_price + result = self.calculate_remaining_budget() / combined_price max_count = max(0, math.floor(result)) if result.is_finite() else items_count return min(items_count, max_count) diff --git a/src/apify/_child_runs.py b/src/apify/_child_runs.py index 82b0d9ac..b6b5e96e 100644 --- a/src/apify/_child_runs.py +++ b/src/apify/_child_runs.py @@ -4,9 +4,10 @@ from collections import defaultdict from contextlib import asynccontextmanager, suppress from dataclasses import dataclass -from datetime import timedelta +from datetime import UTC, datetime, timedelta +from decimal import Decimal from logging import getLogger -from typing import TYPE_CHECKING, Self +from typing import TYPE_CHECKING, Protocol, Self from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError, model_validator from pydantic.alias_generators import to_camel @@ -20,6 +21,7 @@ from apify_client._models import Run from apify_client._resource_clients import RunClientAsync + from apify._charging import ChargingManagerImplementation from apify.storages import KeyValueStore logger = getLogger(__name__) @@ -36,9 +38,26 @@ _ACTIVE_STATUSES = frozenset({'READY', 'RUNNING', 'ABORTING', 'TIMING-OUT'}) +_TERMINAL_STATUSES = frozenset({'SUCCEEDED', 'FAILED', 'ABORTED', 'TIMED-OUT'}) + _STATUS_MAX_AGE = timedelta(seconds=10) """How long an observed active status counts toward the concurrency limit before the run is fetched again.""" +_CHARGE_SETTLE_TIME = timedelta(minutes=3) +"""How long after a run finishes the platform may still add to its `usage_total_usd`.""" + + +class StartRun(Protocol): + """Starts a new run of the Actor with the given charge limit.""" + + def __call__(self, *, max_total_charge_usd: Decimal | None) -> Awaitable[Run]: ... + + +class ResurrectRun(Protocol): + """Resurrects the recorded run, given its run client, with the given charge limit.""" + + def __call__(self, run_client: RunClientAsync, *, max_total_charge_usd: Decimal | None) -> Awaitable[Run]: ... + class ChildRunRecord(BaseModel): """A child run tracked under a name in the child run registry.""" @@ -60,6 +79,15 @@ class ChildRunRecord(BaseModel): abort_with_parent: bool = False """Whether the current run is aborted when this Actor run is gracefully aborted.""" + max_total_charge_usd: Decimal | None = None + """Charge limit of the current run reserved from this Actor run's budget, or `None` when nothing is reserved.""" + + charged_usd: Decimal | None = None + """Final charge of the current run, set once it finished and its `usage_total_usd` settled.""" + + previous_charged_usd: Decimal = Decimal(0) + """Charges of the earlier runs under this name, still counted against this Actor run's budget.""" + @model_validator(mode='after') def _check_started_from(self) -> Self: if (self.actor_id is None) == (self.task_id is None): @@ -94,6 +122,9 @@ class ChildRunInfo: abort_with_parent: bool """Whether the current run is aborted when this Actor run is gracefully aborted.""" + max_total_charge_usd: Decimal | None + """Charge limit of the current run reserved from this Actor run's budget, or `None` when nothing is reserved.""" + _records_adapter = TypeAdapter(dict[str, ChildRunRecord]) @@ -105,8 +136,13 @@ class ChildRunRegistry: starting the child and that write can still orphan the child, since nothing but the platform knows about it. """ - def __init__(self, open_key_value_store: Callable[[], Awaitable[KeyValueStore]]) -> None: + def __init__( + self, + open_key_value_store: Callable[[], Awaitable[KeyValueStore]], + get_charging_manager: Callable[[], ChargingManagerImplementation] | None = None, + ) -> None: self._open_key_value_store = open_key_value_store + self._get_charging_manager = get_charging_manager self._records: dict[str, ChildRunRecord] | None = None self._load_lock = asyncio.Lock() self._write_lock = asyncio.Lock() @@ -120,6 +156,12 @@ def __init__(self, open_key_value_store: Callable[[], Awaitable[KeyValueStore]]) self._observed: dict[str, tuple[str, float]] = {} """Last observed status of the run recorded under each name, with the event loop time it was observed at.""" self._parent_aborting = False + self._reserving: dict[str, Decimal] = {} + """Charge limits reserved for starts and resurrections in flight, not recorded yet.""" + self._unsettled_charges: dict[str, Decimal] = {} + """Charge of each finished current run whose `usage_total_usd` may still grow, as last observed.""" + self._resurrected_after: dict[str, datetime] = {} + """When the current run under each name finished before it was resurrected, to ignore older snapshots of it.""" def set_max_concurrent_runs(self, max_concurrent_runs: int | None) -> None: """Set how many recorded runs may be active at once, or remove the limit with `None`.""" @@ -134,9 +176,10 @@ async def find_or_start( actor_id: str | None = None, task_id: str | None = None, client: ApifyClientAsync, - start_run: Callable[[], Awaitable[Run]], - resurrect_run: Callable[[RunClientAsync], Awaitable[Run]], + start_run: StartRun, + resurrect_run: ResurrectRun, abort_with_parent: bool = False, + max_total_charge_usd: Decimal | None = None, ) -> tuple[Run, bool]: """Return the run recorded under `name`, or start one when there is none to reuse. @@ -146,6 +189,10 @@ async def find_or_start( Starting or resurrecting a run waits while the concurrency limit is reached. Reattaching never waits. + When this Actor run has a `max_total_charge_usd` set by the user, a started or resurrected run gets at most + the part of it that is not charged yet nor reserved for other child runs, and that part stays reserved for the + run until it finishes. A reattached run keeps the limit it was started with. + Args: 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`. @@ -155,6 +202,7 @@ async def find_or_start( 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`. + max_total_charge_usd: Charge limit for a started or resurrected run, lowered to the budget left. Returns: The run, and whether it was newly started. @@ -173,7 +221,10 @@ async def find_or_start( self._clients[name] = client if record is None: - async with self._slot(name, client): + async with ( + self._slot(name, client), + self._budget(name, client, max_total_charge_usd) as (limit, reserved), + ): run = await self._start( name, actor_id=actor_id, @@ -181,6 +232,9 @@ async def find_or_start( start_run=start_run, previous_run_ids=[], abort_with_parent=abort_with_parent, + max_total_charge_usd=limit, + reserved_usd=reserved, + previous_charged_usd=Decimal(0), ) return run, True @@ -190,8 +244,15 @@ async def find_or_start( if run is not None and run.status in _SETTLING_STATUSES: run = await run_client.wait_for_finish() + if run is not None: + await self._settle_charge(name, run) + record = records[name] + if run is None or run.status == 'FAILED': - async with self._slot(name, client): + async with ( + self._slot(name, client), + self._budget(name, client, max_total_charge_usd) as (limit, reserved), + ): run = await self._start( name, actor_id=actor_id, @@ -199,6 +260,9 @@ async def find_or_start( start_run=start_run, previous_run_ids=[*record.previous_run_ids, record.run_id], abort_with_parent=abort_with_parent, + max_total_charge_usd=limit, + reserved_usd=reserved, + previous_charged_usd=record.previous_charged_usd + self._current_charge(name, record), ) return run, True @@ -208,9 +272,19 @@ async def find_or_start( await self._save(name, record.model_copy(update={'abort_with_parent': abort_with_parent})) if run.status in _RESURRECTABLE_STATUSES: - async with self._slot(name, client): + async with ( + self._slot(name, client), + self._budget(name, client, max_total_charge_usd, replaces_current=True) as (limit, reserved), + ): logger.info(f'Resurrecting child run "{name}"', extra={'run_id': run.id, 'status': run.status}) - run = await resurrect_run(run_client) + finished_at = run.finished_at + run = await resurrect_run(run_client, max_total_charge_usd=limit) + if finished_at is not None: + self._resurrected_after[name] = finished_at + self._unsettled_charges.pop(name, None) + await self._save( + name, records[name].model_copy(update={'max_total_charge_usd': reserved, 'charged_usd': None}) + ) self._observe(name, run) return run, False @@ -224,6 +298,7 @@ async def run_finished(self, name: str, run: Run) -> None: if record is None or record.run_id != run.id: return self._observe(name, run) + await self._settle_charge(name, run) async with self._slots: self._slots.notify_all() @@ -241,6 +316,7 @@ async def list_runs(self, client: ApifyClientAsync) -> dict[str, ChildRunInfo]: for name, run in zip(records, runs, strict=True): if run is not None and name in current and current[name].run_id == run.id: self._observe(name, run) + await self._settle_charge(name, run) return { name: ChildRunInfo( actor_id=record.actor_id, @@ -249,6 +325,7 @@ async def list_runs(self, client: ApifyClientAsync) -> dict[str, ChildRunInfo]: run=run, previous_run_ids=list(record.previous_run_ids), abort_with_parent=record.abort_with_parent, + max_total_charge_usd=record.max_total_charge_usd, ) for (name, record), run in zip(records.items(), runs, strict=True) } @@ -291,18 +368,24 @@ async def _start( *, actor_id: str | None, task_id: str | None, - start_run: Callable[[], Awaitable[Run]], + start_run: StartRun, previous_run_ids: list[str], abort_with_parent: bool, + max_total_charge_usd: Decimal | None, + reserved_usd: Decimal | None, + previous_charged_usd: Decimal, ) -> Run: - run = await start_run() + run = await start_run(max_total_charge_usd=max_total_charge_usd) 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, + max_total_charge_usd=reserved_usd, + previous_charged_usd=previous_charged_usd, ) + self._unsettled_charges.pop(name, None) await self._save(name, record) self._observe(name, run) return run @@ -373,6 +456,136 @@ async def _count_active(self, client: ApifyClientAsync, *, exclude: str) -> int: def _observe(self, name: str, run: Run) -> None: self._observed[name] = (run.status, asyncio.get_running_loop().time()) + def reserved_usd(self) -> Decimal: + """Return the part of this Actor run's budget reserved for or charged by its named child runs.""" + records = self._records or {} + return sum( + (record.previous_charged_usd + self._current_charge(name, record) for name, record in records.items()), + start=sum(self._reserving.values(), start=Decimal(0)), + ) + + async def load(self) -> None: + """Load the records persisted by an earlier attempt of this Actor run, so their reservations count.""" + await self._load() + + def _current_charge(self, name: str, record: ChildRunRecord) -> Decimal: + """Return the charge of the current run under `name`, or its whole limit while it may still grow.""" + if record.charged_usd is not None: + return record.charged_usd + if name in self._unsettled_charges: + return self._unsettled_charges[name] + return record.max_total_charge_usd or Decimal(0) + + async def _settle_charge(self, name: str, run: Run) -> None: + """Release the unused part of the limit of a finished current run, recording its charge once it settled.""" + record = (await self._load()).get(name) + if ( + record is None + or record.run_id != run.id + or record.max_total_charge_usd is None + or record.charged_usd is not None + or run.status not in _TERMINAL_STATUSES + or run.usage_total_usd is None + # A snapshot fetched before a resurrection shows the run as it finished the previous time. + or ( + name in self._resurrected_after + and run.finished_at is not None + and run.finished_at <= self._resurrected_after[name] + ) + ): + return + + charged_usd = Decimal(str(run.usage_total_usd)) + if run.finished_at is not None and datetime.now(UTC) - run.finished_at >= _CHARGE_SETTLE_TIME: + await self._save(name, record.model_copy(update={'charged_usd': charged_usd}), if_unchanged=record) + self._unsettled_charges.pop(name, None) + else: + self._unsettled_charges[name] = charged_usd + + @asynccontextmanager + async def _budget( + self, + name: str, + client: ApifyClientAsync, + max_total_charge_usd: Decimal | None, + *, + replaces_current: bool = False, + ) -> AsyncIterator[tuple[Decimal | None, Decimal | None]]: + """Reserve a charge limit for starting or resurrecting the run under `name`, capped at the budget left. + + Yields the limit to start the run with, and the part of it reserved from this Actor run's budget. Nothing is + reserved when this Actor run has no budget set by the user. + + Args: + name: Name of the child run. + client: Client used to fetch recorded runs whose charge is not settled. + max_total_charge_usd: The requested limit, or `None` for all of the budget left. + replaces_current: Whether the new limit replaces the one of the current run under `name`, as a + resurrection does, so that one's reservation is available to it. + """ + charging_manager = self._get_charging_manager() if self._get_charging_manager else None + if charging_manager is None or not await charging_manager.is_max_total_charge_usd_set_by_user(): + yield max_total_charge_usd, None + return + + await self._refresh_charges(client, exclude=name) + async with charging_manager.charge_lock(): + available = charging_manager.calculate_remaining_budget() + record = (await self._load()).get(name) + current_charge = self._current_charge(name, record) if replaces_current and record is not None else 0 + available += current_charge + if available <= 0: + raise RuntimeError( + f'Child run "{name}" was not started, since the budget of this Actor run is spent or reserved for ' + 'other child runs.' + ) + limit = available if max_total_charge_usd is None else min(max_total_charge_usd, available) + if max_total_charge_usd is not None and limit < max_total_charge_usd: + logger.info( + f'Lowering the charge limit of child run "{name}" to {limit} USD, the budget left for it', + extra={'requested_usd': str(max_total_charge_usd)}, + ) + # A resurrected run's current charge is reserved by its record already. + self._reserving[name] = max(limit - current_charge, Decimal(0)) + + try: + yield limit, limit + finally: + self._reserving.pop(name, None) + + async def _refresh_charges(self, client: ApifyClientAsync, *, exclude: str) -> None: + """Fetch recorded runs whose charge is not settled, releasing the unused limit of those that finished.""" + records = await self._load() + now = asyncio.get_running_loop().time() + names = [ + name + for name, record in records.items() + if name != exclude + and name not in self._reserving + and record.max_total_charge_usd is not None + and record.charged_usd is None + # A run seen active a moment ago still holds its whole limit. + and not ( + name in self._observed + and self._observed[name][0] in _ACTIVE_STATUSES + and now - self._observed[name][1] < _STATUS_MAX_AGE.total_seconds() + ) + ] + runs = await asyncio.gather( + *(self._clients.get(name, client).run(records[name].run_id).get() for name in names), + return_exceptions=True, + ) + for name, run in zip(names, runs, strict=True): + if isinstance(run, BaseException): + logger.warning( + f'Failed to fetch child run "{name}" to release its unused budget', + extra={'run_id': records[name].run_id}, + exc_info=run, + ) + elif run is not None: + self._observe(name, run) + await self._settle_charge(name, run) + async def _load(self) -> dict[str, ChildRunRecord]: async with self._load_lock: if self._records is None: @@ -387,10 +600,21 @@ async def _load(self) -> dict[str, ChildRunRecord]: ) from exc return self._records - async def _save(self, name: str, record: ChildRunRecord) -> None: + async def _save(self, name: str, record: ChildRunRecord, *, if_unchanged: ChildRunRecord | None = None) -> None: + """Record `record` under `name`. + + With `if_unchanged`, the write is skipped when `name` no longer holds that record, and a reservation of a start + or resurrection in flight is left alone. + """ records = await self._load() key_value_store = await self._open_key_value_store() async with self._write_lock: + if if_unchanged is not None: + if records.get(name) is not if_unchanged: + return + else: + # The record carries the limit reserved for a start or resurrection in flight from here on. + self._reserving.pop(name, None) records[name] = record await key_value_store.set_value( CHILD_RUNS_KEY, _records_adapter.dump_python(records, by_alias=True, mode='json') diff --git a/tests/e2e/test_actor_child_runs.py b/tests/e2e/test_actor_child_runs.py index 859f3216..db1feeac 100644 --- a/tests/e2e/test_actor_child_runs.py +++ b/tests/e2e/test_actor_child_runs.py @@ -2,6 +2,7 @@ import asyncio from datetime import timedelta +from decimal import Decimal from typing import TYPE_CHECKING from apify import Actor @@ -167,3 +168,40 @@ async def main() -> None: assert run_result.status == 'SUCCEEDED' # The parent run and its two child runs. assert (await actor.runs().list()).total == 3 + + +async def test_named_child_runs_share_the_parent_budget( + make_actor: MakeActorFunction, + run_actor: RunActorFunction, +) -> None: + """Named child runs of a parent started with `max_total_charge_usd` get charge limits within its budget.""" + + async def main() -> None: + from decimal import Decimal + + 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 '' + first = await Actor.start( + actor_id=actor_id, run_input={'is_child': True}, run_name='first', max_total_charge_usd=Decimal('0.25') + ) + second = await Actor.start(actor_id=actor_id, run_input={'is_child': True}, run_name='second') + try: + limits = [] + for run in (first, second): + fetched = await Actor.apify_client.run(run.id).get() + assert fetched is not None, 'fetched is None' + limits.append(fetched.options.max_total_charge_usd) + assert limits == [0.25, 0.75], f'limits={limits}' + finally: + for run in (first, second): + await Actor.apify_client.run(run.id).abort() + + actor = await make_actor(label='child-run-budget', main_func=main) + run_result = await run_actor(actor, max_total_charge_usd=Decimal(1)) + + assert run_result.status == 'SUCCEEDED' diff --git a/tests/unit/actor/test_actor_child_runs.py b/tests/unit/actor/test_actor_child_runs.py index ed5ab2a4..674113c7 100644 --- a/tests/unit/actor/test_actor_child_runs.py +++ b/tests/unit/actor/test_actor_child_runs.py @@ -1,7 +1,8 @@ from __future__ import annotations import asyncio -from datetime import timedelta +from datetime import UTC, datetime, timedelta +from decimal import Decimal from typing import TYPE_CHECKING, Any from unittest.mock import AsyncMock, MagicMock, Mock @@ -13,6 +14,7 @@ from apify import Actor, Configuration from apify._actor import _ActorType +from apify._charging import ChargingManagerImplementation from apify._child_runs import CHILD_RUNS_KEY, ChildRunRegistry from apify.events import ApifyEventManager @@ -76,6 +78,9 @@ async def test_named_start_records_run_in_kvs(apify_client_async_patcher: ApifyC 'runId': 'new-run', 'previousRunIds': [], 'abortWithParent': False, + 'maxTotalChargeUsd': None, + 'chargedUsd': None, + 'previousChargedUsd': '0', } } @@ -196,6 +201,9 @@ async def test_named_start_replaces_failed_or_missing_run( 'runId': 'new-run', 'previousRunIds': ['old-run'], 'abortWithParent': False, + 'maxTotalChargeUsd': None, + 'chargedUsd': None, + 'previousChargedUsd': '0', } } @@ -715,7 +723,7 @@ async def test_aborting_waits_for_a_named_start_in_flight() -> None: started = asyncio.Event() release = asyncio.Event() - async def start_run() -> Run: + async def start_run(*, max_total_charge_usd: Decimal | None) -> Run: # noqa: ARG001 started.set() await release.wait() return make_run('new-run', 'READY') @@ -783,7 +791,7 @@ def make_client(statuses: dict[str, str]) -> Mock: def run(run_id: str) -> Mock: run_client = Mock() run_client.get = AsyncMock(side_effect=lambda: make_run(run_id, statuses[run_id])) - run_client.resurrect = AsyncMock(side_effect=lambda: make_run(run_id, 'RUNNING')) + run_client.resurrect = AsyncMock(side_effect=lambda **_: make_run(run_id, 'RUNNING')) run_client.abort = AsyncMock() return run_client @@ -796,7 +804,7 @@ async def start_child( ) -> Run: """Start a named child run with the registry, adding its run to `statuses` as `RUNNING`.""" - async def start_run() -> Run: + async def start_run(*, max_total_charge_usd: Decimal | None) -> Run: # noqa: ARG001 new_run_id = run_id or f'{name}-run' statuses[new_run_id] = 'RUNNING' return make_run(new_run_id, 'READY') @@ -806,7 +814,9 @@ async def start_run() -> Run: actor_id='some-actor', client=client, start_run=start_run, - resurrect_run=lambda run_client: run_client.resurrect(), + resurrect_run=lambda run_client, max_total_charge_usd: run_client.resurrect( + max_total_charge_usd=max_total_charge_usd + ), ) return run @@ -1084,3 +1094,423 @@ async def test_named_call_task_frees_its_slot_when_the_run_finishes( await asyncio.wait_for(Actor.start('some-actor', run_name='second'), timeout=1) assert len(apify_client_async_patcher.calls['actor']['start']) == 1 + + +@pytest.fixture +def parent_budget( + monkeypatch: pytest.MonkeyPatch, apify_client_async_patcher: ApifyClientAsyncPatcher +) -> dict[str, Run]: + """Give the Actor run a budget of 10 USD, with each local charge costing 1 USD, and a client serving `runs`.""" + monkeypatch.setenv('ACTOR_MAX_TOTAL_CHARGE_USD', '10') + monkeypatch.setenv('ACTOR_TEST_PAY_PER_EVENT', 'true') + runs: dict[str, Run] = {} + + def start(*_args: Any, **_kwargs: Any) -> Run: + run = make_run(f'run-{len(runs) + 1}', 'READY') + runs[run.id] = run.model_copy(update={'status': 'RUNNING'}) + return run + + apify_client_async_patcher.patch('actor', 'start', replacement_method=start) + apify_client_async_patcher.patch( + 'run', 'get', replacement_method=lambda run_client: runs.get(run_client._resource_id) + ) + apify_client_async_patcher.patch( + 'run', 'resurrect', replacement_method=lambda run_client, **_: runs[run_client._resource_id] + ) + return runs + + +def finish(run: Run, status: str, usage_total_usd: float, *, finished_ago: timedelta = timedelta(0)) -> Run: + return run.model_copy( + update={ + 'status': status, + 'usage_total_usd': usage_total_usd, + 'finished_at': datetime.now(UTC) - finished_ago, + } + ) + + +def started_limits(apify_client_async_patcher: ApifyClientAsyncPatcher) -> list[Decimal | None]: + return [kwargs['max_total_charge_usd'] for _, kwargs in apify_client_async_patcher.calls['actor']['start']] + + +@pytest.mark.usefixtures('parent_budget') +async def test_named_start_gets_the_budget_left(apify_client_async_patcher: ApifyClientAsyncPatcher) -> None: + """A named start without a charge limit gets the part of the parent's budget it has not charged itself.""" + async with Actor: + await Actor.charge('some-event', count=3) + await Actor.start('some-actor', run_name='child') + + assert started_limits(apify_client_async_patcher) == [Decimal(7)] + + +@pytest.mark.parametrize( + ('requested', 'expected'), + [ + pytest.param(Decimal(4), Decimal(4), id='within budget'), + pytest.param(Decimal(20), Decimal(10), id='above budget'), + ], +) +@pytest.mark.usefixtures('parent_budget') +async def test_named_start_charge_limit_is_capped_at_the_budget_left( + apify_client_async_patcher: ApifyClientAsyncPatcher, + requested: Decimal, + expected: Decimal, +) -> None: + """An explicit charge limit of a named start is kept within the budget left and lowered above it.""" + async with Actor: + await Actor.start('some-actor', run_name='child', max_total_charge_usd=requested) + + assert started_limits(apify_client_async_patcher) == [expected] + + +@pytest.mark.usefixtures('parent_budget') +async def test_unnamed_start_charge_limit_is_passed_through( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """A start without a name is not tracked, so its charge limit is neither capped nor reserved.""" + async with Actor: + await Actor.start('some-actor', max_total_charge_usd=Decimal(20)) + await Actor.start('some-actor', run_name='child') + + assert started_limits(apify_client_async_patcher) == [Decimal(20), Decimal(10)] + + +async def test_named_start_charge_limit_is_passed_through_without_a_parent_budget( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """Without a parent budget, a named start gets the charge limit it asked for.""" + apify_client_async_patcher.patch('actor', 'start', return_value=make_run('new-run', 'READY')) + + async with Actor: + await Actor.start('some-actor', run_name='first', max_total_charge_usd=Decimal(20)) + await Actor.start('some-actor', run_name='second') + + assert started_limits(apify_client_async_patcher) == [Decimal(20), None] + + +@pytest.mark.usefixtures('parent_budget') +async def test_running_child_run_reserves_its_charge_limit(apify_client_async_patcher: ApifyClientAsyncPatcher) -> None: + """The limit of a running child run is reserved, both from later child runs and from the parent's own charges.""" + async with Actor: + await Actor.start('some-actor', run_name='first', max_total_charge_usd=Decimal(6)) + charge_result = await Actor.charge('some-event', count=3) + await Actor.start('some-actor', run_name='second') + + assert charge_result.charged_count == 3 + assert started_limits(apify_client_async_patcher) == [Decimal(6), Decimal(1)] + + +@pytest.mark.usefixtures('parent_budget') +async def test_named_call_task_is_rejected_under_a_parent_budget( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """A named task call that would start a run under a parent budget set by the user raises without starting it.""" + apify_client_async_patcher.patch('task', 'start', return_value=make_run('task-run', 'READY')) + + async with Actor: + with pytest.raises(RuntimeError, match='cannot be given a part of the budget'): + await Actor.call_task('some-task', run_name='child') + + assert apify_client_async_patcher.calls['task']['start'] == [] + + +async def test_named_call_task_starts_without_a_charge_limit( + apify_client_async_patcher: ApifyClientAsyncPatcher, +) -> None: + """Without a parent budget, a named task call starts the task with the options the task client accepts.""" + apify_client_async_patcher.patch('task', 'start', return_value=make_run('task-run', 'READY')) + apify_client_async_patcher.patch('run', 'wait_for_finish', return_value=make_run('task-run', 'SUCCEEDED')) + + async with Actor: + await Actor.call_task('some-task', run_name='child') + + [(_, kwargs)] = apify_client_async_patcher.calls['task']['start'] + assert 'max_total_charge_usd' not in kwargs + + +@pytest.mark.usefixtures('parent_budget') +async def test_reserved_budget_limits_the_parent_charges() -> None: + """The parent charges only the part of its budget not reserved for child runs.""" + async with Actor: + await Actor.start('some-actor', run_name='child', max_total_charge_usd=Decimal(6)) + charge_result = await Actor.charge('some-event', count=10) + + assert charge_result.charged_count == 4 + + +@pytest.mark.usefixtures('parent_budget') +async def test_exhausted_budget_rejects_a_named_start() -> None: + """A named start raises when the whole parent budget is charged or reserved.""" + async with Actor: + await Actor.start('some-actor', run_name='first') + with pytest.raises(RuntimeError, match='budget of this Actor run is spent or reserved'): + await Actor.start('some-actor', run_name='second') + + +@pytest.mark.usefixtures('parent_budget') +async def test_concurrent_named_starts_share_the_budget() -> None: + """Concurrent named starts reserve their limits one at a time, so together they stay within the budget.""" + async with Actor: + results = await asyncio.gather( + *(Actor.start('some-actor', run_name=name, max_total_charge_usd=Decimal(6)) for name in ('a', 'b')), + return_exceptions=True, + ) + child_runs = await Actor.child_runs() + + assert not any(isinstance(result, BaseException) for result in results) + assert sorted(info.max_total_charge_usd or Decimal(0) for info in child_runs.values()) == [Decimal(4), Decimal(6)] + + +async def test_finished_child_run_releases_its_unused_budget( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher, monkeypatch: pytest.MonkeyPatch +) -> None: + """A finished child run keeps only its charge reserved, and records it once it can no longer change.""" + monkeypatch.setattr('apify._child_runs._STATUS_MAX_AGE', timedelta(0)) + async with Actor: + first = await Actor.start('some-actor', run_name='first', max_total_charge_usd=Decimal(6)) + parent_budget[first.id] = finish(parent_budget[first.id], 'SUCCEEDED', 2) + await Actor.start('some-actor', run_name='second', max_total_charge_usd=Decimal(3)) + parent_budget['run-2'] = finish(parent_budget['run-2'], 'SUCCEEDED', 1, finished_ago=timedelta(minutes=5)) + await Actor.start('some-actor', run_name='third') + kvs = await Actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert started_limits(apify_client_async_patcher) == [Decimal(6), Decimal(3), Decimal(7)] + # The first run finished just now, so the platform may still add to its charge. + assert stored['first']['chargedUsd'] is None + assert stored['second']['chargedUsd'] == '1' + + +async def seed_budget_record(name: str, run_id: str, **fields: Any) -> None: + """Seed the registry with a record carrying budget fields, as an earlier attempt of this Actor run would.""" + kvs = await Actor.open_key_value_store() + await kvs.set_value( + CHILD_RUNS_KEY, {name: {'actorId': 'some-actor', 'runId': run_id, 'previousRunIds': [], **fields}} + ) + + +@pytest.mark.usefixtures('parent_budget') +async def test_reservations_of_an_earlier_attempt_limit_the_parent_charges() -> None: + """A child run recorded by an earlier attempt of the parent keeps its limit reserved from the parent's charges.""" + async with Actor: + await seed_budget_record('child', 'old-run', maxTotalChargeUsd='6') + + # A fresh instance, as the parent is after a migration or resurrection. + async with _ActorType() as actor: + charge_result = await actor.charge('some-event', count=10) + + assert charge_result.charged_count == 4 + + +async def test_resurrection_reuses_the_reservation_of_its_run( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """A resurrected run gets the budget left plus its own reservation, since its limit covers its earlier charges.""" + parent_budget['old-run'] = finish(make_run('old-run', 'RUNNING'), 'ABORTED', 1) + parent_budget['other-run'] = make_run('other-run', 'RUNNING') + + async with Actor: + kvs = await Actor.open_key_value_store() + await kvs.set_value( + CHILD_RUNS_KEY, + { + 'child': {'actorId': 'some-actor', 'runId': 'old-run', 'maxTotalChargeUsd': '6'}, + 'other': {'actorId': 'some-actor', 'runId': 'other-run', 'maxTotalChargeUsd': '3'}, + }, + ) + + async with _ActorType() as actor: + await actor.start('some-actor', run_name='child') + + [(_, kwargs)] = apify_client_async_patcher.calls['run']['resurrect'] + assert kwargs['max_total_charge_usd'] == Decimal(7) + + +async def test_replaced_failed_run_keeps_its_charge_reserved( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """A failed run replaced by a new one under the same name keeps its charge counted against the budget.""" + parent_budget['old-run'] = finish(make_run('old-run', 'RUNNING'), 'FAILED', 3, finished_ago=timedelta(minutes=5)) + + async with Actor: + await seed_budget_record('child', 'old-run', maxTotalChargeUsd='6') + + async with _ActorType() as actor: + await actor.start('some-actor', run_name='child') + kvs = await actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert started_limits(apify_client_async_patcher) == [Decimal(7)] + assert stored['child']['previousChargedUsd'] == '3' + assert stored['child']['maxTotalChargeUsd'] == '7' + + +@pytest.mark.usefixtures('parent_budget') +async def test_failed_named_start_releases_its_reservation(apify_client_async_patcher: ApifyClientAsyncPatcher) -> None: + """A named start that fails leaves no part of the budget reserved.""" + async with Actor: + apify_client_async_patcher.patch('actor', 'start', replacement_method=Mock(side_effect=RuntimeError('boom'))) + with pytest.raises(RuntimeError, match='boom'): + await Actor.start('some-actor', run_name='child') + charge_result = await Actor.charge('some-event', count=10) + + assert charge_result.charged_count == 10 + + +async def test_named_start_in_flight_reserves_its_limit( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """The limit of a named start in flight is reserved before the platform returns its run.""" + started = asyncio.Event() + release = asyncio.Event() + + async def start(*_args: Any, **_kwargs: Any) -> Run: + started.set() + await release.wait() + run = make_run('slow-run', 'READY') + parent_budget[run.id] = run + return run + + async with Actor: + apify_client_async_patcher.patch('actor', 'start', replacement_method=start) + start_task = asyncio.create_task(Actor.start('some-actor', run_name='first')) + await started.wait() + charge_result = await Actor.charge('some-event', count=1) + release.set() + await start_task + + assert charge_result.charged_count == 0 + + +async def test_listing_child_runs_releases_the_unused_budget(parent_budget: dict[str, Run]) -> None: + """Listing the child runs releases the unused limit of those that finished.""" + async with Actor: + run = await Actor.start('some-actor', run_name='child', max_total_charge_usd=Decimal(6)) + parent_budget[run.id] = finish(parent_budget[run.id], 'SUCCEEDED', 2) + await Actor.child_runs() + charge_result = await Actor.charge('some-event', count=10) + + assert charge_result.charged_count == 8 + + +async def test_named_call_releases_the_unused_budget_when_the_run_finishes( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """A named call releases the unused limit of its run once the run finishes.""" + apify_client_async_patcher.patch( + 'run', + 'wait_for_finish', + replacement_method=lambda run_client, **_: finish(parent_budget[run_client._resource_id], 'SUCCEEDED', 2), + ) + + async with Actor: + await Actor.call('some-actor', run_name='child', max_total_charge_usd=Decimal(6), logger=None) + charge_result = await Actor.charge('some-event', count=10) + + assert charge_result.charged_count == 8 + + +@pytest.mark.usefixtures('parent_budget') +async def test_platform_default_charge_limit_is_not_shared_with_child_runs( + apify_client_async_patcher: ApifyClientAsyncPatcher, monkeypatch: pytest.MonkeyPatch +) -> None: + """A limit the platform gave the parent by default is not split among its child runs.""" + monkeypatch.setattr( + ChargingManagerImplementation, 'is_max_total_charge_usd_set_by_user', AsyncMock(return_value=False) + ) + + async with Actor: + await Actor.start('some-actor', run_name='first') + await Actor.start('some-actor', run_name='second', max_total_charge_usd=Decimal(20)) + charge_result = await Actor.charge('some-event', count=10) + + assert started_limits(apify_client_async_patcher) == [None, Decimal(20)] + assert charge_result.charged_count == 10 + + +async def test_run_fetched_before_its_resurrection_keeps_the_reservation( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """A snapshot of a run fetched before its resurrection does not release the limit of the resurrected run.""" + parent_budget['old-run'] = finish(make_run('old-run', 'RUNNING'), 'ABORTED', 1, finished_ago=timedelta(minutes=5)) + listing = asyncio.Event() + release = asyncio.Event() + + async def get(run_client: Any) -> Run | None: + run = parent_budget.get(run_client._resource_id) + if not listing.is_set(): + listing.set() + await release.wait() + return run + + def resurrect(run_client: Any, **_: Any) -> Run: + run = parent_budget[run_client._resource_id].model_copy(update={'status': 'RUNNING', 'finished_at': None}) + parent_budget[run.id] = run + return run + + async with Actor: + await seed_budget_record('child', 'old-run', maxTotalChargeUsd='6') + + apify_client_async_patcher.patch('run', 'get', replacement_method=get, is_async=True) + apify_client_async_patcher.patch('run', 'resurrect', replacement_method=resurrect, is_async=True) + + async with _ActorType() as actor: + list_task = asyncio.create_task(actor.child_runs()) + await listing.wait() + await actor.start('some-actor', run_name='child') + release.set() + await list_task + charge_result = await actor.charge('some-event', count=10) + + assert charge_result.charged_count == 0 + + +async def test_resurrection_in_flight_reserves_its_limit_once( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher +) -> None: + """While a resurrection is in flight, its limit is reserved once, including the charge its run made before.""" + parent_budget['old-run'] = finish(make_run('old-run', 'RUNNING'), 'ABORTED', 1, finished_ago=timedelta(minutes=5)) + started = asyncio.Event() + release = asyncio.Event() + + async def resurrect(run_client: Any, **_: Any) -> Run: + started.set() + await release.wait() + return parent_budget[run_client._resource_id].model_copy(update={'status': 'RUNNING', 'finished_at': None}) + + async with Actor: + await seed_budget_record('child', 'old-run', maxTotalChargeUsd='6') + + apify_client_async_patcher.patch('run', 'resurrect', replacement_method=resurrect, is_async=True) + + async with _ActorType() as actor: + start_task = asyncio.create_task(actor.start('some-actor', run_name='child', max_total_charge_usd=Decimal(4))) + await started.wait() + charge_result = await actor.charge('some-event', count=10) + release.set() + await start_task + + assert charge_result.charged_count == 6 + + +async def test_failed_resurrection_leaves_the_charge_of_its_run_to_settle( + parent_budget: dict[str, Run], apify_client_async_patcher: ApifyClientAsyncPatcher, monkeypatch: pytest.MonkeyPatch +) -> None: + """A resurrection that fails does not stop the charge of the finished run from being recorded later.""" + parent_budget['old-run'] = finish(make_run('old-run', 'RUNNING'), 'ABORTED', 1) + + async with Actor: + await seed_budget_record('child', 'old-run', maxTotalChargeUsd='6') + + apify_client_async_patcher.patch('run', 'resurrect', replacement_method=Mock(side_effect=RuntimeError('boom'))) + + async with _ActorType() as actor: + with pytest.raises(RuntimeError, match='boom'): + await actor.start('some-actor', run_name='child') + monkeypatch.setattr('apify._child_runs._CHARGE_SETTLE_TIME', timedelta(0)) + await actor.child_runs() + kvs = await actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert stored['child']['chargedUsd'] == '1' diff --git a/tests/unit/actor/test_charging_manager.py b/tests/unit/actor/test_charging_manager.py index f0ace2e9..16f28f39 100644 --- a/tests/unit/actor/test_charging_manager.py +++ b/tests/unit/actor/test_charging_manager.py @@ -689,3 +689,51 @@ async def test_charge_registers_the_count_capped_by_the_budget(mock_client: Magi assert (await cm.charge('search', count=5, idempotency_key='key-1')).charged_count == 2 assert cm.get_charged_event_count('search') == 2 assert mock_client.run.return_value.charge.await_count == 1 + + +@pytest.mark.parametrize( + ('options_extra', 'expected'), + [ + pytest.param({'isMaxTotalChargeUsdSetByUser': True}, True, id='set by user'), + pytest.param({'isMaxTotalChargeUsdSetByUser': False}, False, id='platform default'), + pytest.param({}, False, id='not reported'), + ], +) +async def test_max_total_charge_usd_set_by_user_is_read_from_the_run_options( + mock_client: MagicMock, *, options_extra: dict[str, Any], expected: bool +) -> None: + """On the platform, whether the limit was set by the user comes from the run options, fetched once.""" + run = MagicMock() + run.options.model_extra = options_extra + mock_client.run.return_value.get = AsyncMock(return_value=run) + config = _make_config( + is_at_home=True, + actor_run_id='run-id', + actor_pricing_info=_make_ppe_pricing_info(), + charged_event_counts={}, + max_total_charge_usd=Decimal(10), + ) + cm = ChargingManagerImplementation(config, mock_client) + async with cm: + assert await cm.is_max_total_charge_usd_set_by_user() is expected + assert await cm.is_max_total_charge_usd_set_by_user() is expected + + mock_client.run.return_value.get.assert_awaited_once() + + +@pytest.mark.parametrize( + ('max_total_charge_usd', 'expected'), + [ + pytest.param(Decimal(10), True, id='limited'), + pytest.param(None, False, id='unlimited'), + ], +) +async def test_max_total_charge_usd_set_by_user_locally( + mock_client: MagicMock, *, max_total_charge_usd: Decimal | None, expected: bool +) -> None: + """Locally, any limit counts as set by the user, and no run is fetched.""" + cm = ChargingManagerImplementation(_make_config(max_total_charge_usd=max_total_charge_usd), mock_client) + async with cm: + assert await cm.is_max_total_charge_usd_set_by_user() is expected + + mock_client.run.return_value.get.assert_not_awaited()