From ef1ba5328ef1f64224389dbed863235d3a936d0c Mon Sep 17 00:00:00 2001 From: Vlada Dusek Date: Fri, 9 Oct 2026 17:05:02 +0200 Subject: [PATCH] feat: share the parent's charge budget with named child runs --- src/apify/_actor.py | 46 ++- src/apify/_charging.py | 34 +- src/apify/_child_runs.py | 247 ++++++++++++- tests/e2e/test_actor_child_runs.py | 38 ++ tests/unit/actor/test_actor_child_runs.py | 414 +++++++++++++++++++++- tests/unit/actor/test_charging_manager.py | 48 +++ 6 files changed, 790 insertions(+), 37 deletions(-) diff --git a/src/apify/_actor.py b/src/apify/_actor.py index 2106379db..bd4b65f5d 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 ChildRunRegistry +from apify._child_runs import 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 @@ -163,7 +163,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).""" @@ -223,6 +225,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: @@ -1027,7 +1030,11 @@ async def start( max_items: Deprecated, use `max_total_charge_usd` instead. Will be removed in version 5.0.0. Works only with legacy pay-per-result Actors. max_total_charge_usd: A limit on the total charged amount, in USD. Once the run exceeds it, the platform - aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. + aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. 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 @@ -1068,7 +1075,6 @@ async def start( content_type=content_type, build=build, max_items=max_items, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=self._resolve_run_timeout(timeout), @@ -1077,7 +1083,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, @@ -1222,7 +1228,11 @@ async def call( max_items: Deprecated, use `max_total_charge_usd` instead. Will be removed in version 5.0.0. Works only with legacy pay-per-result Actors. max_total_charge_usd: A limit on the total charged amount, in USD. Once the run exceeds it, the platform - aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. + aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. 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 @@ -1289,7 +1299,6 @@ async def call( content_type=content_type, build=build, max_items=max_items, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=self._resolve_run_timeout(timeout), @@ -1322,7 +1331,7 @@ async def _find_or_start_child_run( task_id: str | None = None, run_input: Any, client: ApifyClientAsync, - start_run: Callable[[], Awaitable[Run]], + start_run: StartRun, build: str | None, max_items: int | None, max_total_charge_usd: Decimal | None, @@ -1338,7 +1347,7 @@ async def _find_or_start_child_run( run_input=run_input, client=client, start_run=start_run, - resurrect_run=lambda run_id: client.run(run_id).resurrect( + resurrect_run=lambda run_id, *, max_total_charge_usd: client.run(run_id).resurrect( build=build, memory_mbytes=memory_mbytes, run_timeout=self._resolve_run_timeout(timeout), @@ -1347,6 +1356,7 @@ async def _find_or_start_child_run( restart_on_error=restart_on_error, ), abort_with_parent=abort_with_parent, + max_total_charge_usd=max_total_charge_usd, ) def _remove_internal_listeners(self) -> None: @@ -1417,7 +1427,11 @@ async def start_task( max_items: Deprecated, use `max_total_charge_usd` instead. Will be removed in version 5.0.0. Works only with legacy pay-per-result Actors. max_total_charge_usd: A limit on the total charged amount, in USD. Once the run exceeds it, the platform - aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. + aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. 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 Task 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 @@ -1454,7 +1468,6 @@ async def start_task( task_input=task_input, build=build, max_items=max_items, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=self._resolve_run_timeout(timeout), @@ -1462,7 +1475,7 @@ async def start_task( ) 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, @@ -1514,7 +1527,11 @@ async def call_task( max_items: Deprecated, use `max_total_charge_usd` instead. Will be removed in version 5.0.0. Works only with legacy pay-per-result Actors. max_total_charge_usd: A limit on the total charged amount, in USD. Once the run exceeds it, the platform - aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. + aborts the run, which takes a few seconds, so the final charge can slightly overshoot the limit. 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 Task 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 @@ -1572,7 +1589,6 @@ async def call_task( task_input=task_input, build=build, max_items=max_items, - max_total_charge_usd=max_total_charge_usd, restart_on_error=restart_on_error, memory_mbytes=memory_mbytes, run_timeout=self._resolve_run_timeout(timeout), diff --git a/src/apify/_charging.py b/src/apify/_charging.py index aebc1c1f5..8bb7cd947 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 5ba335af0..105e3439c 100644 --- a/src/apify/_child_runs.py +++ b/src/apify/_child_runs.py @@ -4,9 +4,10 @@ import hashlib import json from contextlib import asynccontextmanager, suppress -from datetime import datetime, timedelta +from datetime import UTC, datetime, timedelta +from decimal import Decimal from logging import getLogger -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Protocol from weakref import WeakValueDictionary from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError @@ -19,6 +20,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__) @@ -35,15 +37,32 @@ _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`.""" + _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.""" _NOT_FOUND_RETRY_INTERVAL_SECS = 0.25 +class StartRun(Protocol): + """Starts a new run of the Actor or task 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 ID, with the given charge limit.""" + + def __call__(self, run_id: str, *, max_total_charge_usd: Decimal | None) -> Awaitable[Run]: ... + + class ChildRunSnapshot(BaseModel): """A child run as last observed by this Actor run.""" @@ -71,6 +90,15 @@ class ChildRunRecord(ChildRunSnapshot): 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.""" + 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.""" @@ -108,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._lock = asyncio.Lock() """Guards loading the records and writing them back to the key-value store.""" @@ -124,6 +157,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`.""" @@ -139,9 +178,10 @@ async def find_or_start( task_id: str | None = None, run_input: Any, client: ApifyClientAsync, - start_run: Callable[[], Awaitable[Run]], - resurrect_run: Callable[[str], 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. @@ -151,6 +191,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. @@ -161,6 +205,7 @@ async def find_or_start( 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`. + 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. @@ -178,9 +223,19 @@ async def find_or_start( ) 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, checksum=checksum, start_run=start_run, history=[], abort_with_parent=abort_with_parent + name, + checksum=checksum, + start_run=start_run, + history=[], + abort_with_parent=abort_with_parent, + max_total_charge_usd=limit, + reserved_usd=reserved, + previous_charged_usd=Decimal(0), ) self._clients[name] = client return run, True @@ -192,19 +247,29 @@ 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': replaced = ChildRunSnapshot( run_id=record.run_id, status=run.status if run is not None else 'LOST', started_at=record.started_at, ) - 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, checksum=checksum, start_run=start_run, history=[*record.history, replaced], 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 @@ -214,9 +279,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.id) + finished_at = run.finished_at + run = await resurrect_run(run.id, 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) else: logger.info(f'Reattaching to child run "{name}"', extra={'run_id': run.id, 'status': run.status}) @@ -253,7 +328,8 @@ async def update(self, name: str, run: Run) -> None: return self._observe(name, run) if record.status != run.status: - await self._save(name, record.model_copy(update={'status': run.status})) + await self._save(name, record.model_copy(update={'status': run.status}), if_unchanged=record) + await self._settle_charge(name, run) async with self._slots: self._slots.notify_all() @@ -294,11 +370,14 @@ async def _start( name: str, *, checksum: str, - start_run: Callable[[], Awaitable[Run]], + start_run: StartRun, history: list[ChildRunSnapshot], 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( run_id=run.id, status=run.status, @@ -306,7 +385,10 @@ async def _start( checksum=checksum, history=history, 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 @@ -377,6 +459,132 @@ 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)), + ) + + 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]: """Read the records from the default key-value store, replacing any read before.""" async with self._lock: @@ -394,10 +602,21 @@ async def load(self) -> dict[str, ChildRunRecord]: async def _load(self) -> dict[str, ChildRunRecord]: return self._records if self._records is not None else await self.load() - 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._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 a9eefb3b5..45811a05c 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 @@ -164,3 +165,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 564de66a3..174dd4c43 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, _child_runs from apify._actor import _ActorType +from apify._charging import ChargingManagerImplementation from apify._child_runs import CHILD_RUNS_KEY, ChildRunRegistry, checksum_request from apify.events import ApifyEventManager @@ -64,6 +66,9 @@ def stored_record( 'checksum': checksum_request(actor_id=actor_id, task_id=task_id, run_input=run_input), 'history': history or [], 'abortWithParent': abort_with_parent, + 'maxTotalChargeUsd': None, + 'chargedUsd': None, + 'previousChargedUsd': '0', } @@ -936,7 +941,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') @@ -1005,7 +1010,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 @@ -1018,7 +1023,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') @@ -1029,7 +1034,9 @@ async def start_run() -> Run: run_input=None, client=client, start_run=start_run, - resurrect_run=lambda run_id: client.run(run_id).resurrect(), + resurrect_run=lambda run_id, *, max_total_charge_usd: client.run(run_id).resurrect( + max_total_charge_usd=max_total_charge_usd + ), ) return run @@ -1277,3 +1284,400 @@ 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_gets_the_budget_left(apify_client_async_patcher: ApifyClientAsyncPatcher) -> None: + """A named task call gets the part of the parent's budget it has not charged itself, as an Actor call does.""" + 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.charge('some-event', count=3) + await Actor.call_task('some-task', run_name='child') + + [(_, kwargs)] = apify_client_async_patcher.calls['task']['start'] + assert kwargs['max_total_charge_usd'] == Decimal(7) + + +@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, + ) + kvs = await Actor.open_key_value_store() + stored = await kvs.get_value(CHILD_RUNS_KEY) + + assert not any(isinstance(result, BaseException) for result in results) + assert sorted(Decimal(record['maxTotalChargeUsd']) for record in stored.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: {**stored_record(run_id, 'RUNNING'), **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': {**stored_record('old-run', 'RUNNING'), 'maxTotalChargeUsd': '6'}, + 'other': {**stored_record('other-run', 'RUNNING'), '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_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: + # Another named start fetches the run under `child` to release its unused budget, and is held mid-fetch. + other_task = asyncio.create_task(actor.start('some-actor', run_name='other')) + await listing.wait() + await actor.start('some-actor', run_name='child') + release.set() + with pytest.raises(RuntimeError, match='spent or reserved'): + await other_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)) + # Another named start fetches the finished run under `child` to release its unused budget. + await actor.start('some-actor', run_name='other') + 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 f0ace2e94..16f28f392 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()