From ce360ed8aea4ec09659f10d25b79f9f61745d638 Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 05:41:59 -0400 Subject: [PATCH 1/8] refactor(python): type durable authoring boundaries --- maintainers/quality-policy.lock.json | 28 +- pyproject.toml | 22 +- scripts/check_quality_policy.py | 2 +- src/volcano_sdk/durable_authoring.py | 319 ++++++++++++++++------- tests/typing/durable_authoring.py | 9 +- tests/unit/fixtures/durable_context.py | 108 ++++++++ tests/unit/fixtures/durable_engine.py | 31 +++ tests/unit/fixtures/invalid_callbacks.py | 6 + tests/unit/test_durable_authoring.py | 309 +++++++++++++++++----- uv.lock | 2 +- 10 files changed, 672 insertions(+), 164 deletions(-) create mode 100644 tests/unit/fixtures/durable_context.py create mode 100644 tests/unit/fixtures/durable_engine.py diff --git a/maintainers/quality-policy.lock.json b/maintainers/quality-policy.lock.json index d8568201..59e90f6b 100644 --- a/maintainers/quality-policy.lock.json +++ b/maintainers/quality-policy.lock.json @@ -5,6 +5,15 @@ "src/volcano_sdk/_generated" ], "executionEnvironments": [ + { + "reportExplicitAny": "error", + "root": "src/volcano_sdk/durable_authoring.py" + }, + { + "reportAny": "error", + "reportExplicitAny": "error", + "root": "tests/unit/test_durable_authoring.py" + }, { "reportAny": "error", "reportExplicitAny": "error", @@ -335,6 +344,15 @@ "$MYPY_CONFIG_FILE_DIR/features", "$MYPY_CONFIG_FILE_DIR/typings" ], + "overrides": [ + { + "disallow_any_explicit": true, + "module": [ + "volcano_sdk.durable_authoring", + "test_durable_authoring" + ] + } + ], "python_version": "3.11", "strict": true, "strict_bytes": true, @@ -548,12 +566,18 @@ "python", "-I", "-c", - "from importlib.metadata import version; assert version('typing-extensions') == '4.10.0'; import volcano_sdk" + "from importlib.metadata import version; assert version('typing-extensions') == '4.12.2'" + ], + [ + "python", + "-I", + "tests/package/optional_dependency.py", + "base" ] ], "constrain_package_deps": true, "deps": [ - "typing-extensions==4.10.0" + "typing-extensions==4.12.2" ] }, "package-types": { diff --git a/pyproject.toml b/pyproject.toml index 84007d16..8515c15f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,7 +24,7 @@ dependencies = [ "attrs>=26.1.0", "centrifuge-python>=0.6.0,<0.7.0", "httpx>=0.28.1,<0.29.0", - "typing-extensions>=4.10.0", + "typing-extensions>=4.12.2", ] [project.optional-dependencies] @@ -132,6 +132,10 @@ mypy_path = [ ] exclude = ["src/volcano_sdk/_generated/"] +[[tool.mypy.overrides]] +module = ["volcano_sdk.durable_authoring", "test_durable_authoring"] +disallow_any_explicit = true + [tool.basedpyright] stubPath = "typings" include = ["src", "tests", "scripts", "features", "typings"] @@ -166,6 +170,15 @@ reportIncompatibleVariableOverride = "error" # Ruff still enforces private access everywhere else. reportPrivateUsage = false +[[tool.basedpyright.executionEnvironments]] +root = "src/volcano_sdk/durable_authoring.py" +reportExplicitAny = "error" + +[[tool.basedpyright.executionEnvironments]] +root = "tests/unit/test_durable_authoring.py" +reportAny = "error" +reportExplicitAny = "error" + [[tool.basedpyright.executionEnvironments]] root = "src/volcano_sdk/storage.py" reportAny = "error" @@ -557,9 +570,12 @@ package = "wheel" commands = [["python", "-I", "tests/package/optional_dependency.py", "base"]] [tool.tox.env.package-min-typing] -deps = ["typing-extensions==4.10.0"] +deps = ["typing-extensions==4.12.2"] constrain_package_deps = true -commands = [["python", "-I", "-c", "from importlib.metadata import version; assert version('typing-extensions') == '4.10.0'; import volcano_sdk"]] +commands = [ + ["python", "-I", "-c", "from importlib.metadata import version; assert version('typing-extensions') == '4.12.2'"], + ["python", "-I", "tests/package/optional_dependency.py", "base"], +] [tool.tox.env.package-durable] extras = ["durable"] diff --git a/scripts/check_quality_policy.py b/scripts/check_quality_policy.py index d9a7a411..b11cb503 100644 --- a/scripts/check_quality_policy.py +++ b/scripts/check_quality_policy.py @@ -18,7 +18,7 @@ from collections.abc import Iterable GENERATED = "src/volcano_sdk/_generated" -LOCK_SHA256 = "5da168f90e02c72a3ee45f4f7b8899149a80c408574c2b693c3d83db7b5d7c1b" +LOCK_SHA256 = "4a6e8f9b39cf3a86ae00517ee5cbe957375009bffab8834be4ea20107bad5817" TYPE_FIXTURES = { "tests/typing/contract_steps.py", "tests/typing/durable_callbacks.py", diff --git a/src/volcano_sdk/durable_authoring.py b/src/volcano_sdk/durable_authoring.py index 2e9dbf30..b67f47df 100644 --- a/src/volcano_sdk/durable_authoring.py +++ b/src/volcano_sdk/durable_authoring.py @@ -27,23 +27,25 @@ def handler(event, ctx): import functools import importlib -from dataclasses import dataclass +from collections.abc import Callable +from dataclasses import dataclass, replace from typing import ( TYPE_CHECKING, - Any, Generic, + Never, ParamSpec, Protocol, TypeAlias, - TypeVar, - cast, + TypeGuard, overload, ) +from typing_extensions import TypeVar + from ._callbacks import require_callable if TYPE_CHECKING: - from collections.abc import Callable, Mapping, Sequence + from collections.abc import Mapping, Sequence from aws_durable_execution_sdk_python.config import ( CompletionConfig, @@ -57,7 +59,12 @@ def handler(event, ctx): RetryStrategyConfig, ) -T = TypeVar("T") +T = TypeVar("T", default=object) +U = TypeVar("U", default=object) +# An unsubscripted public handler can accept a specific input type without +# promising that callers may pass an arbitrary object to it. +_Input = TypeVar("_Input", default=Never) +_Output = TypeVar("_Output", default=object) _P = ParamSpec("_P") # A duration: "30s", "5m", "2h", "1d", a compound string like "1m30s", a whole # number of seconds, or the mapping form. @@ -65,7 +72,7 @@ def handler(event, ctx): # What a durable function is written as, and what the platform invokes it as. # The second argument differs: the handler is given a durable context, and the # wrapper is given the invocation's own context. -DurableHandler: TypeAlias = "Callable[[Any, DurableContext], Any]" +DurableHandler: TypeAlias = Callable[[_Input, "DurableContext"], _Output] # Invocation envelopes are distinct from the user handler's input and result. FunctionHandler: TypeAlias = "Callable[[object, object], object]" @@ -104,9 +111,15 @@ def handler(event, ctx): _INVALID_BRANCH = "a parallel branch is a callable, or a ParallelBranch" _INVALID_ITEMS = "map() requires a sequence of items" _INVALID_WAIT_ARGS = "wait() takes a name and a duration, or a duration alone" + + # Distinguishes an omitted initial_state from an explicit None, which is a # legitimate state for a condition to start from. -_UNSET: Any = object() +class _Unset: + __slots__ = () + + +_UNSET = _Unset() class DurableRuntimeMissingError(Exception): @@ -138,17 +151,86 @@ def seconds(self, value: int) -> object: ... def step_options(self, *, retry: Retry, at_most_once: bool) -> object: ... - def wait_condition_options(self, options: WaitUntilOptions) -> object: ... + def wait_condition_options(self, options: WaitUntilOptions[T]) -> object: ... def map_options(self, options: BatchOptions | None) -> object: ... def parallel_options(self, options: BatchOptions | None) -> object: ... def named_branch( - self, run: Callable[[object], object], name: str | None + self, run: Callable[[_RuntimeContext], object], name: str | None ) -> object: ... +class _OperationScope(Protocol): + logger: DurableLogger + attempt: int + + +class _RuntimeBatchItem(Protocol[T]): + index: int + status: object + result: T | None + error: object + + +class _RuntimeBatch(Protocol[T]): + success_count: int + failure_count: int + completion_reason: object + + def succeeded(self) -> list[_RuntimeBatchItem[T]]: ... + + def failed(self) -> list[_RuntimeBatchItem[T]]: ... + + def get_results(self) -> list[T]: ... + + def get_errors(self) -> list[object]: ... + + def throw_if_error(self) -> None: ... + + +class _RuntimeContext(Protocol): + logger: DurableLogger + + def step( + self, + func: Callable[[_OperationScope], T], + name: str | None, + config: object, + ) -> T: ... + + def wait(self, duration: object, name: str | None = None) -> None: ... + + def run_in_child_context( + self, func: Callable[[_RuntimeContext], T], name: str | None + ) -> T: ... + + def wait_for_condition( + self, + func: Callable[[T, _OperationScope], T], + config: object, + name: str | None, + ) -> T: ... + + def map( + self, + items: list[U], + func: Callable[[_RuntimeContext, U, int, list[U]], T], + name: str | None, + config: object, + ) -> _RuntimeBatch[T]: ... + + def parallel( + self, + branches: list[Callable[[_RuntimeContext], T] | object], + name: str | None, + config: object, + ) -> _RuntimeBatch[T]: ... + + def set_logger(self, logger: object) -> None: ... + + class _Engine: """The durable protocol, from the AWS durable execution SDK. @@ -174,7 +256,7 @@ def __init__(self) -> None: waits = importlib.import_module(f"{_ENGINE_MODULE}.waits") root = importlib.import_module(_ENGINE_MODULE) except ImportError as error: - raise DurableRuntimeMissingError(error) from error + raise DurableRuntimeMissingError from error self.durable_execution = root.durable_execution self.duration: type[EngineDuration] = config.Duration self.step_config: type[StepConfig] = config.StepConfig @@ -183,7 +265,9 @@ def __init__(self) -> None: self.parallel_config: type[ParallelConfig] = config.ParallelConfig self.completion_config: type[CompletionConfig] = config.CompletionConfig self.parallel_branch = config.ParallelBranch - self.create_retry_strategy = retries.create_retry_strategy + self.create_retry_strategy: Callable[ + [RetryStrategyConfig], Callable[[Exception, int], RetryDecision] + ] = retries.create_retry_strategy self.retry_strategy_config: type[RetryStrategyConfig] = ( retries.RetryStrategyConfig ) @@ -220,26 +304,40 @@ def step_options(self, *, retry: Retry, at_most_once: bool) -> StepConfig: The runtime step configuration. """ - config: dict[str, Any] = {} + config = self.step_config() if at_most_once: - config["step_semantics"] = self.step_semantics.AT_MOST_ONCE_PER_RETRY + config = replace( + config, step_semantics=self.step_semantics.AT_MOST_ONCE_PER_RETRY + ) strategy = self._retry_strategy(retry) if strategy is not None: - config["retry_strategy"] = strategy - return self.step_config(**config) + config = replace(config, retry_strategy=strategy) + return config - def _retry_strategy(self, retry: object) -> Any: + def _retry_strategy( + self, retry: object + ) -> Callable[[Exception, int], RetryDecision] | None: if retry is None or retry is True: return None # False disables the runtime's default retry policy for this step. if retry is False: return self._never_retry() + return self._custom_retry_strategy(retry) + + def _custom_retry_strategy( + self, retry: object + ) -> Callable[[Exception, int], RetryDecision]: if isinstance(retry, RetryOptions): - return self.create_retry_strategy( - self.retry_strategy_config(**self._retry_kwargs(retry)) - ) + return self.create_retry_strategy(self._retry_config(retry)) if callable(retry): - return retry + + def decide(error: Exception, attempt: int) -> RetryDecision: + result = retry(error, attempt) + if not isinstance(result, self.retry_decision): + raise TypeError(_INVALID_RETRY) + return result + + return decide raise TypeError(_INVALID_RETRY) def _never_retry(self) -> Callable[[Exception, int], RetryDecision]: @@ -250,24 +348,39 @@ def never_retry(_error: Exception, _attempt: int) -> RetryDecision: return never_retry - def _retry_kwargs(self, retry: RetryOptions) -> dict[str, Any]: - return _engine_kwargs( - max_attempts=retry.attempts, - initial_delay=self._optional_duration(retry.initial_delay, "initial_delay"), - max_delay=self._optional_duration(retry.max_delay, "max_delay"), - backoff_rate=retry.backoff_rate, - retryable_errors=(None if retry.retry_on is None else list(retry.retry_on)), - retryable_error_types=( - None if retry.retry_on_types is None else list(retry.retry_on_types) - ), - ) + def _retry_config(self, retry: RetryOptions) -> RetryStrategyConfig: + config = self.retry_strategy_config() + self._set_retry_timing(config, retry) + self._set_retry_filters(config, retry) + return config + + def _set_retry_timing( + self, config: RetryStrategyConfig, retry: RetryOptions + ) -> None: + if retry.attempts is not None: + config.max_attempts = retry.attempts + if retry.initial_delay is not None: + config.initial_delay = self.seconds( + _to_seconds(retry.initial_delay, "initial_delay") + ) + if retry.max_delay is not None: + config.max_delay = self.seconds(_to_seconds(retry.max_delay, "max_delay")) + if retry.backoff_rate is not None: + config.backoff_rate = retry.backoff_rate + + @staticmethod + def _set_retry_filters(config: RetryStrategyConfig, retry: RetryOptions) -> None: + if retry.retry_on is not None: + config.retryable_errors = list(retry.retry_on) + if retry.retry_on_types is not None: + config.retryable_error_types = list(retry.retry_on_types) def _optional_duration( self, value: Duration | None, field_name: str ) -> EngineDuration | None: return None if value is None else self.seconds(_to_seconds(value, field_name)) - def wait_condition_options(self, options: WaitUntilOptions) -> object: + def wait_condition_options(self, options: WaitUntilOptions[T]) -> object: """Build the runtime's polling options. Returns: @@ -276,7 +389,7 @@ def wait_condition_options(self, options: WaitUntilOptions) -> object: """ until = options.until - def keep_polling(state: Any) -> bool: + def keep_polling(state: T) -> bool: return not until(state) strategy = self.wait_strategy_config( @@ -293,15 +406,6 @@ def keep_polling(state: Any) -> bool: initial_state=options.initial_state, ) - def _batch_options(self, options: BatchOptions | None) -> dict[str, Any]: - resolved = BatchOptions() if options is None else options - config = _engine_kwargs(max_concurrency=resolved.concurrency) - if resolved.min_succeeded is not None: - config["completion_config"] = self.completion_config( - min_successful=resolved.min_succeeded - ) - return config - def map_options(self, options: BatchOptions | None) -> object: """Build the runtime's map options. @@ -309,7 +413,18 @@ def map_options(self, options: BatchOptions | None) -> object: The runtime map configuration. """ - return self.map_config(**self._batch_options(options)) + resolved = BatchOptions() if options is None else options + config = self.map_config() + if resolved.concurrency is not None: + config = replace(config, max_concurrency=resolved.concurrency) + if resolved.min_succeeded is not None: + config = replace( + config, + completion_config=self.completion_config( + min_successful=resolved.min_succeeded + ), + ) + return config def parallel_options(self, options: BatchOptions | None) -> ParallelConfig: """Build the runtime's parallel options. @@ -318,9 +433,22 @@ def parallel_options(self, options: BatchOptions | None) -> ParallelConfig: The runtime parallel configuration. """ - return self.parallel_config(**self._batch_options(options)) + resolved = BatchOptions() if options is None else options + config = self.parallel_config() + if resolved.concurrency is not None: + config = replace(config, max_concurrency=resolved.concurrency) + if resolved.min_succeeded is not None: + config = replace( + config, + completion_config=self.completion_config( + min_successful=resolved.min_succeeded + ), + ) + return config - def named_branch(self, run: Callable[[object], object], name: str | None) -> object: + def named_branch( + self, run: Callable[[_RuntimeContext], object], name: str | None + ) -> object: """Build a named runtime branch. Returns: @@ -350,21 +478,23 @@ class RetryOptions: retry_on_types: Sequence[type[Exception]] | None = None -Retry: TypeAlias = "bool | RetryOptions | Callable[[Exception, int], Any] | None" +Retry: TypeAlias = ( + "bool | RetryOptions | Callable[[Exception, int], RetryDecision] | None" +) @dataclass(frozen=True, slots=True) -class WaitUntilOptions: +class WaitUntilOptions(Generic[T]): """How `wait_until` polls, and what it polls for.""" # Stop waiting once this returns true for the state the check returned. - until: Callable[[Any], bool] + until: Callable[[T], bool] # The state a check receives. Required, and distinguished from an explicit # None: the wait starts by asking `until` about it. Treat it as the state # every check starts from rather than an accumulator -- a check should # decide from what it observes now, because the platform does not promise # to carry a previous check's return into the next one. - initial_state: Any = _UNSET + initial_state: T | _Unset = _UNSET # Delay before the second check, then multiplied by backoff_rate up to # max_interval. interval: Duration | None = None @@ -501,7 +631,7 @@ class BatchResult(Generic[T]): "succeeded", ) - def __init__(self, batch: Any) -> None: + def __init__(self, batch: _RuntimeBatch[T]) -> None: """Flatten an engine batch result.""" self._batch = batch self.items: tuple[BatchItem[T], ...] = tuple( @@ -531,7 +661,7 @@ def throw_if_failed(self) -> None: self._batch.throw_if_error() -def _batch_items(batch: Any) -> list[Any]: +def _batch_items(batch: _RuntimeBatch[T]) -> list[_RuntimeBatchItem[T]]: """List the items that finished, in input order. The in-flight ones are left out on purpose: see `BatchResult`. @@ -544,7 +674,7 @@ def _batch_items(batch: Any) -> list[Any]: return sorted(items, key=lambda item: item.index) -def _batch_failure(error: Any) -> BatchFailure | None: +def _batch_failure(error: object) -> BatchFailure | None: if error is None: return None return BatchFailure( @@ -554,7 +684,7 @@ def _batch_failure(error: Any) -> BatchFailure | None: ) -def _completion_reason(batch: Any) -> str | None: +def _completion_reason(batch: object) -> str | None: reason = getattr(batch, "completion_reason", None) if reason is None: return None @@ -579,7 +709,7 @@ class DurableContext: __slots__ = ("_context", "_engine", "log") - def __init__(self, context: Any, engine: _DurableEngine) -> None: + def __init__(self, context: _RuntimeContext, engine: _DurableEngine) -> None: """Wrap an engine context.""" self._context = context self._engine = engine @@ -611,18 +741,13 @@ def step( """ step_name, step_func = _named(name, func, "step") - def run(scope: Any) -> T: + def run(scope: _OperationScope) -> T: return step_func(StepScope(scope.logger, scope.attempt)) - # The engine hands back whatever the step returned, untyped. T is the - # wrapper's own contract with the caller. - return cast( - "T", - self._context.step( - run, - step_name, - self._engine.step_options(retry=retry, at_most_once=at_most_once), - ), + return self._context.step( + run, + step_name, + self._engine.step_options(retry=retry, at_most_once=at_most_once), ) def wait(self, name: str | Duration, duration: Duration | None = None) -> None: @@ -658,17 +783,17 @@ def child( child_name, child_func = _named(name, func, "child") engine = self._engine - def run(context: Any) -> T: + def run(context: _RuntimeContext) -> T: return child_func(DurableContext(context, engine)) - return cast("T", self._context.run_in_child_context(run, child_name)) + return self._context.run_in_child_context(run, child_name) def wait_until( self, - check: Callable[[Any, StepScope], Any], - options: WaitUntilOptions, + check: Callable[[T, StepScope], T], + options: WaitUntilOptions[T], name: str | None = None, - ) -> Any: + ) -> T: """Poll until a condition holds, suspending between checks. The check reports the current state and `options.until` decides @@ -687,7 +812,7 @@ def wait_until( _validate_wait_options(options) engine = self._engine - def check_state(state: Any, scope: Any) -> Any: + def check_state(state: T, scope: _OperationScope) -> T: return check_func(state, StepScope(scope.logger, scope.attempt)) return self._context.wait_for_condition( @@ -698,8 +823,8 @@ def check_state(state: Any, scope: Any) -> Any: def map( self, - items: Sequence[Any], - func: Callable[[Any, DurableContext, int], T], + items: Sequence[U], + func: Callable[[U, DurableContext, int], T], name: str | None = None, options: BatchOptions | None = None, ) -> BatchResult[T]: @@ -720,7 +845,7 @@ def map( raise TypeError(_INVALID_ITEMS) engine = self._engine - def run(context: Any, item: Any, index: int, _all: Any) -> T: + def run(context: _RuntimeContext, item: U, index: int, _all: list[U]) -> T: return map_func(item, DurableContext(context, engine), index) return BatchResult( @@ -754,7 +879,9 @@ def parallel( ) ) - def _branch(self, branch: Callable[[DurableContext], T] | ParallelBranch[T]) -> Any: + def _branch( + self, branch: Callable[[DurableContext], T] | ParallelBranch[T] + ) -> object: """Adapt one branch to the engine. The branch's own result type is not carried through: what goes to the @@ -770,18 +897,16 @@ def _branch(self, branch: Callable[[DurableContext], T] | ParallelBranch[T]) -> """ engine = self._engine if isinstance(branch, ParallelBranch): - # isinstance cannot carry the branch's type argument, so its - # function is read at the erased type the engine takes anyway. - named = cast("Callable[[DurableContext], Any]", branch.run) + named: Callable[[DurableContext], T] = branch.run - def run_named(context: Any) -> Any: + def run_named(context: _RuntimeContext) -> T: return named(DurableContext(context, engine)) return engine.named_branch(run_named, branch.name) require_callable(branch, _INVALID_BRANCH) - bare = branch + bare: Callable[[DurableContext], T] = branch - def run_bare(context: Any) -> Any: + def run_bare(context: _RuntimeContext) -> object: return bare(DurableContext(context, engine)) return run_bare @@ -798,20 +923,22 @@ def _wait_duration(self, value: object) -> object: @overload -def durable(handler: DurableHandler, *, logger: object = ...) -> FunctionHandler: ... +def durable( + handler: DurableHandler[T, U], *, logger: object = ... +) -> FunctionHandler: ... @overload def durable( handler: None = ..., *, logger: object = ... -) -> Callable[[DurableHandler], FunctionHandler]: ... +) -> Callable[[DurableHandler[T, U]], FunctionHandler]: ... def durable( - handler: DurableHandler | None = None, + handler: DurableHandler[T, U] | None = None, *, logger: object = None, -) -> FunctionHandler | Callable[[DurableHandler], FunctionHandler]: +) -> FunctionHandler | Callable[[DurableHandler[T, U]], FunctionHandler]: """Wrap a handler so Volcano runs it as a durable execution. Usable bare or with options:: @@ -834,7 +961,7 @@ def handler(event, ctx): ... """ if handler is None: - def decorate(func: DurableHandler) -> FunctionHandler: + def decorate(func: DurableHandler[T, U]) -> FunctionHandler: return durable(func, logger=logger) return decorate @@ -860,7 +987,7 @@ def invoke(event: object, function_context: object) -> object: if not wrapped: engine = _Engine.load() - def run(input_value: T, context: Any) -> object: + def run(input_value: T, context: _RuntimeContext) -> object: if logger is not None: context.set_logger(logger) return handler(input_value, DurableContext(context, engine)) @@ -871,7 +998,7 @@ def run(input_value: T, context: Any) -> object: return invoke -def _validate_wait_options(options: WaitUntilOptions) -> None: +def _validate_wait_options(options: WaitUntilOptions[T]) -> None: if not callable(options.until): raise TypeError(_REQUIRES_UNTIL) if options.timeout is not None: @@ -906,7 +1033,7 @@ def _callable(func: Callable[_P, T] | None, operation: str) -> Callable[_P, T]: return func -def _engine_kwargs(**entries: Any) -> dict[str, Any]: +def _engine_kwargs(**entries: object) -> dict[str, object]: """Drop what the caller left out. The engine's configs are dataclasses with real defaults, so passing None @@ -935,13 +1062,25 @@ def _to_seconds(value: object, field_name: str) -> int: """ if isinstance(value, (int, float)): return _numeric_seconds(value, field_name) + if _is_string_keyed_mapping(value): + return _mapping_seconds(value, field_name) if isinstance(value, dict): - return _mapping_seconds(cast("dict[str, object]", value), field_name) + raise TypeError(_duration_type_error(field_name)) if not isinstance(value, str): raise TypeError(_duration_type_error(field_name)) return _parse_duration(value.strip(), field_name) +def _is_string_keyed_mapping(value: object) -> TypeGuard[dict[str, object]]: + if not _is_object_dict(value): + return False + return all(isinstance(key, str) for key in value) + + +def _is_object_dict(value: object) -> TypeGuard[dict[object, object]]: + return isinstance(value, dict) + + def _numeric_seconds(value: object, field_name: str) -> int: if isinstance(value, bool): raise TypeError(_duration_type_error(field_name)) diff --git a/tests/typing/durable_authoring.py b/tests/typing/durable_authoring.py index 205e0e15..553fdd3a 100644 --- a/tests/typing/durable_authoring.py +++ b/tests/typing/durable_authoring.py @@ -2,7 +2,12 @@ from typing import TypedDict, assert_type -from volcano_sdk.durable_authoring import DurableContext, FunctionHandler, durable +from volcano_sdk.durable_authoring import ( + DurableContext, + DurableHandler, + FunctionHandler, + durable, +) class Order(TypedDict): @@ -29,6 +34,8 @@ def configured(event: Order, _context: DurableContext) -> int: wrapped = durable(order_quantity) +legacy_handler: DurableHandler = order_quantity +specific_handler: DurableHandler[Order, int] = order_quantity _: object = assert_type(wrapped, FunctionHandler) _ = assert_type(bare, FunctionHandler) _ = assert_type(called, FunctionHandler) diff --git a/tests/unit/fixtures/durable_context.py b/tests/unit/fixtures/durable_context.py new file mode 100644 index 00000000..7ebe9ee5 --- /dev/null +++ b/tests/unit/fixtures/durable_context.py @@ -0,0 +1,108 @@ +"""A typed runtime context that records parallel calls without scheduling them.""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Generic, TypeVar + +if TYPE_CHECKING: + from collections.abc import Callable + + from volcano_sdk.durable_authoring import ( + DurableLogger, + _OperationScope, + _RuntimeBatch, + _RuntimeBatchItem, + _RuntimeContext, + ) + +T = TypeVar("T") +U = TypeVar("U") +_UNEXPECTED_OPERATION = "unexpected runtime operation in parallel test" + + +class EmptyBatch(Generic[T]): + """A completed batch with no items.""" + + success_count = 0 + failure_count = 0 + completion_reason: object = "all_completed" + + def succeeded(self) -> list[_RuntimeBatchItem[T]]: + return [] + + def failed(self) -> list[_RuntimeBatchItem[T]]: + return [] + + def get_results(self) -> list[T]: + return [] + + def get_errors(self) -> list[object]: + return [] + + def throw_if_error(self) -> None: + return + + +class RecordingContext: + """Implement the runtime context and capture its parallel call.""" + + logger: DurableLogger = logging.getLogger(__name__) + + def __init__(self) -> None: + self.branches: list[object] | None = None + self.name: str | None = None + self.config: object = None + + def step( + self, + func: Callable[[_OperationScope], T], + name: str | None, + config: object, + ) -> T: + _ = (func, name, config) + raise AssertionError(_UNEXPECTED_OPERATION) + + def wait(self, duration: object, name: str | None = None) -> None: + _ = (duration, name) + raise AssertionError(_UNEXPECTED_OPERATION) + + def run_in_child_context( + self, func: Callable[[_RuntimeContext], T], name: str | None + ) -> T: + _ = (func, name) + raise AssertionError(_UNEXPECTED_OPERATION) + + def wait_for_condition( + self, + func: Callable[[T, _OperationScope], T], + config: object, + name: str | None, + ) -> T: + _ = (func, config, name) + raise AssertionError(_UNEXPECTED_OPERATION) + + def map( + self, + items: list[U], + func: Callable[[_RuntimeContext, U, int, list[U]], T], + name: str | None, + config: object, + ) -> _RuntimeBatch[T]: + _ = (items, func, name, config) + raise AssertionError(_UNEXPECTED_OPERATION) + + def parallel( + self, + branches: list[Callable[[_RuntimeContext], T] | object], + name: str | None, + config: object, + ) -> _RuntimeBatch[T]: + self.branches = list(branches) + self.name = name + self.config = config + return EmptyBatch[T]() + + def set_logger(self, logger: object) -> None: + _ = logger + raise AssertionError(_UNEXPECTED_OPERATION) diff --git a/tests/unit/fixtures/durable_engine.py b/tests/unit/fixtures/durable_engine.py new file mode 100644 index 00000000..554de651 --- /dev/null +++ b/tests/unit/fixtures/durable_engine.py @@ -0,0 +1,31 @@ +"""Check the optional runtime's dynamic import boundary before scheduling work.""" + +from __future__ import annotations + +import importlib + +from volcano_sdk import durable_authoring + + +def assert_runtime_surface() -> None: + """Fail promptly if the installed runtime lacks an adapter dependency.""" + engine = durable_authoring._Engine.load() + root = importlib.import_module("aws_durable_execution_sdk_python") + config = importlib.import_module("aws_durable_execution_sdk_python.config") + retries = importlib.import_module("aws_durable_execution_sdk_python.retries") + waits = importlib.import_module("aws_durable_execution_sdk_python.waits") + + assert engine.durable_execution is root.durable_execution + assert engine.duration is config.Duration + assert engine.step_config is config.StepConfig + assert engine.step_semantics is config.StepSemantics + assert engine.map_config is config.MapConfig + assert engine.parallel_config is config.ParallelConfig + assert engine.completion_config is config.CompletionConfig + assert engine.parallel_branch is config.ParallelBranch + assert engine.create_retry_strategy is retries.create_retry_strategy + assert engine.retry_strategy_config is retries.RetryStrategyConfig + assert engine.retry_decision is retries.RetryDecision + assert engine.create_wait_strategy is waits.create_wait_strategy + assert engine.wait_strategy_config is waits.WaitStrategyConfig + assert engine.wait_for_condition_config is waits.WaitForConditionConfig diff --git a/tests/unit/fixtures/invalid_callbacks.py b/tests/unit/fixtures/invalid_callbacks.py index 3d53a3ae..f6b67fae 100644 --- a/tests/unit/fixtures/invalid_callbacks.py +++ b/tests/unit/fixtures/invalid_callbacks.py @@ -23,5 +23,11 @@ def register_non_callable_branch(context: DurableContext) -> None: _ = context.parallel(["not a branch"]) # type: ignore[list-item] +def run_non_callable_operation(context: DurableContext, operation: str) -> object: + if operation == "step": + return context.step("named", "not a function") # type: ignore[arg-type] + return context.child("named", "not a function") # type: ignore[arg-type] + + def use_non_callable_retry(context: DurableContext) -> None: context.step("charge", lambda _scope: None, retry="aggressively") # type: ignore[arg-type] diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index f5ac9b5f..4043354f 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -10,17 +10,28 @@ import inspect import json import logging +from collections.abc import Mapping from contextlib import contextmanager from types import ModuleType, SimpleNamespace -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, TypeGuard import pytest -from aws_durable_execution_sdk_python.config import Duration +from aws_durable_execution_sdk_python.config import ( + Duration, + MapConfig, + ParallelConfig, +) +from aws_durable_execution_sdk_python.config import ( + ParallelBranch as EngineParallelBranch, +) from aws_durable_execution_sdk_python.retries import RetryDecision from aws_durable_execution_sdk_python_testing import DurableFunctionTestRunner +from fixtures.durable_context import RecordingContext +from fixtures.durable_engine import assert_runtime_surface from fixtures.invalid_callbacks import ( decorate_non_callable, register_non_callable_branch, + run_non_callable_operation, use_non_callable_retry, ) from fixtures.invalid_wait_options import invalid_wait_duration, non_callable_predicate @@ -58,6 +69,74 @@ _EXPECTED_FAILURE = "expected the execution to fail" +def _is_mapping(value: object) -> TypeGuard[Mapping[object, object]]: + return isinstance(value, Mapping) + + +def _is_object_list(value: object) -> TypeGuard[list[object]]: + return isinstance(value, list) + + +def _ready(state: object) -> bool: + if not _is_mapping(state): + msg = "expected a state mapping" + raise TypeError(msg) + ready: object = state.get("ready") + if not isinstance(ready, bool): + msg = "expected a boolean ready state" + raise TypeError(msg) + return ready + + +def test_an_omitted_mapping_duration_part_is_zero() -> None: + assert to_seconds({"seconds": 5}, "wait") == 5 + + +@pytest.mark.parametrize("value", [{1: 5}, {"seconds": 5, 1: 3}]) +def test_duration_mapping_refuses_non_string_keys(value: object) -> None: + with pytest.raises(TypeError, match="must be a duration string"): + _ = to_seconds(value, "wait") + + +def test_retry_false_produces_an_immediate_no_retry_decision() -> None: + decision = durable_authoring._Engine()._never_retry()(RuntimeError("failed"), 1) + + assert isinstance(decision, RetryDecision) + assert decision.should_retry is False + assert decision.delay.to_seconds() == 0 + + +def test_map_options_forward_both_batch_limits() -> None: + config = durable_authoring._Engine().map_options( + BatchOptions(concurrency=2, min_succeeded=1) + ) + + assert isinstance(config, MapConfig) + assert config.max_concurrency == 2 + assert config.completion_config.min_successful == 1 + + +def test_parallel_forwards_branches_name_and_batch_limits() -> None: + runtime = RecordingContext() + context = DurableContext(runtime, durable_authoring._Engine()) + result = context.parallel( + [ParallelBranch(lambda _child: "named", name="alpha"), lambda _child: "bare"], + "fan-out", + BatchOptions(concurrency=2, min_succeeded=1), + ) + + assert result.completed == 0 + assert runtime.name == "fan-out" + assert isinstance(runtime.config, ParallelConfig) + assert runtime.config.max_concurrency == 2 + assert runtime.config.completion_config.min_successful == 1 + assert runtime.branches is not None + assert len(runtime.branches) == 2 + assert isinstance(runtime.branches[0], EngineParallelBranch) + assert runtime.branches[0].name == "alpha" + assert callable(runtime.branches[1]) + + @contextmanager def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunner]: """Run a handler on the local runner, then close what it leaves open. @@ -70,6 +149,9 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne Yields: The local runner, closed when the context exits. """ + # Fail an unavailable or malformed runtime before entering the scheduler: + # it otherwise waits for a result that the handler cannot produce. + assert_runtime_surface() runner = DurableFunctionTestRunner(handler) try: yield runner @@ -81,7 +163,7 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne loop.close() -def run_handler(handler: Any, event: object = None) -> Any: +def run_handler(handler: FunctionHandler, event: object = None) -> object: """Run a durable handler to completion on the local runner. Returns: @@ -97,7 +179,7 @@ def run_handler(handler: Any, event: object = None) -> Any: return None if result.result is None else json.loads(result.result) -def failing_handler(handler: Any, event: object = None) -> str: +def failing_handler(handler: FunctionHandler, event: object = None) -> str: """Run a handler expected to fail. Returns: @@ -115,7 +197,7 @@ def failing_handler(handler: Any, event: object = None) -> str: def test_a_step_result_is_recorded_and_returned() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.step("charge", lambda scope: {"id": "ch_1", "on": scope.attempt}) assert run_handler(handler, {"order_id": "o-1"}) == {"id": "ch_1", "on": 1} @@ -123,7 +205,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_an_unnamed_step_still_runs() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.step(lambda _scope: "done") assert run_handler(handler) == "done" @@ -131,7 +213,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_the_handler_receives_the_execution_input() -> None: @durable - def handler(event: Any, ctx: DurableContext) -> Any: + def handler(event: dict[str, str], ctx: DurableContext) -> object: return ctx.step("echo", lambda _scope: event["order_id"]) assert run_handler(handler, {"order_id": "o-42"}) == "o-42" @@ -139,7 +221,7 @@ def handler(event: Any, ctx: DurableContext) -> Any: def test_a_wait_suspends_and_resumes_the_execution() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: first = ctx.step("first", lambda _scope: 1) ctx.wait("settle", "1s") return first + ctx.step("second", lambda _scope: 1) @@ -149,7 +231,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_a_wait_takes_a_duration_alone() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: ctx.wait("1s") return "resumed" @@ -158,7 +240,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_operations_are_recorded_under_their_names() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: _ = ctx.step("charge", lambda _scope: "ch_1") ctx.wait("settle", "1s") return ctx.child("fulfil", lambda child: child.step("ship", lambda _s: "ok")) @@ -177,7 +259,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_a_child_context_groups_its_own_operations() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.child( "fulfil", lambda child: { @@ -191,7 +273,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_map_runs_the_work_over_every_item() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: batch = ctx.map( [1, 2, 3], lambda item, child, _index: child.step(lambda _s: item * 10), @@ -218,7 +300,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_map_respects_a_concurrency_limit() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: batch = ctx.map( [1, 2, 3, 4], lambda item, child, _index: child.step(lambda _s: item), @@ -232,7 +314,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_parallel_runs_named_and_bare_branches() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: branches: list[Callable[[DurableContext], str] | ParallelBranch[str]] = [ ParallelBranch(lambda child: child.step(lambda _s: "a"), name="alpha"), lambda child: child.step(lambda _s: "b"), @@ -243,9 +325,17 @@ def handler(_event: Any, ctx: DurableContext) -> Any: assert run_handler(handler) == {"results": ["a", "b"], "completed": 2} +def test_parallel_options_apply_both_batch_limits() -> None: + engine = durable_authoring._Engine.load() + config = engine.parallel_options(BatchOptions(concurrency=2, min_succeeded=1)) + + assert config.max_concurrency == 2 + assert config.completion_config.min_successful == 1 + + def test_a_batch_reports_the_item_that_failed() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def work(item: int, child: DurableContext, _index: int) -> int: def run(_scope: StepScope) -> int: if item == 2: @@ -290,7 +380,7 @@ def test_a_failure_is_reported_in_the_facade_s_own_shape() -> None: """ @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: batch = ctx.map([1, 2], fail_second_item, "one-fails") # Found rather than indexed. `items` carries the items that settled, # and a batch can come back the moment the failure does -- leaving the @@ -328,7 +418,7 @@ def test_an_early_completion_reports_only_what_finished() -> None: """ @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def work(item: int, child: DurableContext, _index: int) -> int: def run(_scope: StepScope) -> int: # The later items outlive the completion threshold, so the batch @@ -347,17 +437,19 @@ def run(_scope: StepScope) -> int: } result = run_handler(handler) - - assert "started" not in result["statuses"], ( + assert _is_mapping(result) + statuses: object = result.get("statuses") + assert _is_object_list(statuses) + assert "started" not in statuses, ( "an in-flight item is not guaranteed to come back on a replay" ) - assert result["completed"] == len(result["statuses"]) + assert result["completed"] == len(statuses) assert result["reason"] == "min_successful_reached" def test_throw_if_failed_surfaces_a_batch_failure() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def work(_item: int, child: DurableContext, _index: int) -> int: def run(_scope: StepScope) -> int: message = "always bad" @@ -373,7 +465,7 @@ def run(_scope: StepScope) -> int: def test_a_step_retries_until_it_succeeds() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def flaky(scope: StepScope) -> int: if scope.attempt < 3: message = f"attempt {scope.attempt} failed" @@ -391,7 +483,7 @@ def flaky(scope: StepScope) -> int: def test_retry_false_fails_on_the_first_error() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def always(scope: StepScope) -> int: message = f"failed on attempt {scope.attempt}" raise RuntimeError(message) @@ -405,7 +497,7 @@ def always(scope: StepScope) -> int: def test_retry_options_stop_after_their_attempt_budget() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def always(scope: StepScope) -> int: message = f"failed on attempt {scope.attempt}" raise RuntimeError(message) @@ -421,7 +513,7 @@ def always(scope: StepScope) -> int: def test_retry_on_limits_which_errors_are_retried() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def always(scope: StepScope) -> int: message = f"unretryable on attempt {scope.attempt}" raise RuntimeError(message) @@ -439,9 +531,39 @@ def always(scope: StepScope) -> int: assert "unretryable on attempt 1" in failing_handler(handler) +def test_retry_config_keeps_unset_defaults_and_sets_requested_fields() -> None: + engine = durable_authoring._Engine.load() + defaults = engine.retry_strategy_config() + config = engine._retry_config( + RetryOptions( + max_delay="9s", + backoff_rate=1.25, + retry_on_types=[ValueError], + ) + ) + + assert config.max_attempts == defaults.max_attempts + assert config.initial_delay == defaults.initial_delay + assert config.max_delay.to_seconds() == 9 + assert config.backoff_rate == pytest.approx(1.25) + assert config.retryable_error_types == [ValueError] + + +def test_custom_retry_must_return_a_runtime_decision() -> None: + engine = durable_authoring._Engine.load() + + def invalid_retry(_error: Exception, _attempt: int) -> str: + return "invalid" + + retry = engine._custom_retry_strategy(invalid_retry) + + with pytest.raises(TypeError, match="retry must be False"): + _ = retry(RuntimeError("failed"), 1) + + def test_a_custom_retry_callable_decides_per_attempt() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: def decide(_error: Exception, attempt: int) -> RetryDecision: return RetryDecision( should_retry=attempt < 2, @@ -461,7 +583,7 @@ def flaky(scope: StepScope) -> int: def test_at_most_once_runs_the_step_once() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.step("charge", lambda _scope: "charged", at_most_once=True) assert run_handler(handler) == "charged" @@ -471,15 +593,17 @@ def test_wait_until_polls_until_the_condition_holds() -> None: polls = {"count": 0} @durable - def handler(_event: Any, ctx: DurableContext) -> Any: - def check(_state: Any, _scope: StepScope) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: + def check( + _state: dict[str, int | bool], _scope: StepScope + ) -> dict[str, int | bool]: polls["count"] += 1 return {"ready": polls["count"] >= 3, "polls": polls["count"]} return ctx.wait_until( check, - WaitUntilOptions( - until=lambda state: bool(state["ready"]), + WaitUntilOptions[dict[str, int | bool]]( + until=_ready, initial_state={"ready": False, "polls": 0}, interval="1s", max_attempts=10, @@ -492,11 +616,11 @@ def check(_state: Any, _scope: StepScope) -> Any: def test_wait_until_fails_when_it_runs_out_of_attempts() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.wait_until( lambda _state, _scope: {"ready": False}, WaitUntilOptions( - until=lambda state: bool(state["ready"]), + until=_ready, initial_state={"ready": False}, interval="1s", max_attempts=2, @@ -547,7 +671,7 @@ def test_a_replacement_logger_is_installed() -> None: recorder = Recorder() @durable(logger=recorder) - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: ctx.log.info("from the handler") return "logged" @@ -557,11 +681,11 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_durable_is_usable_bare_and_called() -> None: @durable - def bare(_event: Any, _ctx: DurableContext) -> Any: + def bare(_event: object, _ctx: DurableContext) -> object: return "bare" @durable() - def called(_event: Any, _ctx: DurableContext) -> Any: + def called(_event: object, _ctx: DurableContext) -> object: return "called" assert run_handler(bare) == "bare" @@ -569,7 +693,7 @@ def called(_event: Any, _ctx: DurableContext) -> Any: def test_the_wrapper_keeps_the_handler_name() -> None: - def order_pipeline(_event: Any, _ctx: DurableContext) -> Any: + def order_pipeline(_event: object, _ctx: DurableContext) -> object: """Handle an order. Returns: @@ -592,16 +716,17 @@ def test_durable_refuses_something_that_is_not_callable() -> None: @pytest.mark.parametrize("operation", ["step", "child"]) def test_an_operation_refuses_a_non_callable(operation: str) -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: - return getattr(ctx, operation)("named", "not a function") + def handler(_event: object, ctx: DurableContext) -> object: + return run_non_callable_operation(ctx, operation) assert f"{operation}() requires a function to run" in failing_handler(handler) def test_wait_refuses_a_name_that_is_not_a_string() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: - return ctx.wait(30, "30s") + def handler(_event: object, ctx: DurableContext) -> object: + ctx.wait(30, "30s") + return None assert "takes a name and a duration" in failing_handler(handler) @@ -615,23 +740,25 @@ def handler(_event: Any, ctx: DurableContext) -> Any: ) def test_wait_refuses_a_wait_of_nothing(duration: object) -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: - return invalid_wait_duration(ctx, duration) + def handler(_event: object, ctx: DurableContext) -> object: + invalid_wait_duration(ctx, duration) + return None assert "wait must be at least 1 second" in failing_handler(handler) def test_wait_refuses_a_wait_longer_than_an_execution_may_run() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: - return ctx.wait("cool-off", {"days": 367}) + def handler(_event: object, ctx: DurableContext) -> object: + ctx.wait("cool-off", {"days": 367}) + return None assert "wait must be at most 31622400 seconds" in failing_handler(handler) def test_wait_until_refuses_a_timeout() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.wait_until( lambda state, _scope: state, WaitUntilOptions(until=bool, initial_state=False, timeout="1h"), @@ -644,10 +771,10 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_wait_until_requires_an_initial_state() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.wait_until( lambda state, _scope: state, - WaitUntilOptions(until=bool), + WaitUntilOptions[object](until=bool), ) assert "requires an `initial_state`" in failing_handler(handler) @@ -670,10 +797,13 @@ def test_missing_batch_completion_reason_remains_absent(batch: object) -> None: def test_wait_until_accepts_none_as_an_initial_state() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: + def check(_state: str | None, _scope: StepScope) -> str | None: + return "ready" + return ctx.wait_until( - lambda _state, _scope: "ready", - WaitUntilOptions( + check, + WaitUntilOptions[str | None]( until=lambda state: bool(state == "ready"), initial_state=None, interval="1s", @@ -688,7 +818,7 @@ def handler(_event: Any, ctx: DurableContext) -> Any: def test_map_refuses_a_string_of_items() -> None: @durable - def handler(_event: Any, ctx: DurableContext) -> Any: + def handler(_event: object, ctx: DurableContext) -> object: return ctx.map("abc", lambda item, _child, _index: item) # A string is a sequence, so mapping over one would otherwise run the work @@ -772,6 +902,50 @@ def test_an_unknown_duration_field_names_what_it_accepts() -> None: _ = to_seconds({"milliseconds": 500}, "interval") +@pytest.mark.parametrize( + ("value", "error_type", "message"), + [ + ( + True, + TypeError, + ( + "interval must be a duration string, a whole number of seconds, " + "or a mapping of days, hours, minutes, seconds" + ), + ), + ( + 1.5, + TypeError, + "interval must be a whole number of seconds, not a fraction", + ), + ( + {"seconds": 1, "years": 2, "milliseconds": 3}, + TypeError, + ( + "interval duration takes days, hours, minutes, seconds " + "(got milliseconds, years)" + ), + ), + ( + {}, + TypeError, + "interval duration needs one of days, hours, minutes, seconds", + ), + ( + {"seconds": -1}, + ValueError, + "interval duration seconds must be a non-negative whole number", + ), + ], +) +def test_invalid_duration_reports_the_field_and_reason( + value: object, error_type: type[Exception], message: str +) -> None: + with pytest.raises(error_type) as raised: + _ = to_seconds(value, "interval") + assert str(raised.value) == message + + @pytest.fixture def without_engine(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: """Hide the durable engine and clear the cached load either side.""" @@ -793,7 +967,7 @@ def blocked(name: str, package: str | None = None) -> ModuleType: @pytest.mark.usefixtures("without_engine") def test_a_missing_runtime_is_reported_on_invocation() -> None: @durable - def handler(_event: Any, _ctx: DurableContext) -> Any: + def handler(_event: object, _ctx: DurableContext) -> object: return None # Decorating has to succeed without the engine: the SDK also runs in @@ -806,21 +980,24 @@ def handler(_event: Any, _ctx: DurableContext) -> Any: @pytest.mark.usefixtures("without_engine") def test_the_missing_runtime_error_says_to_deploy_as_durable() -> None: @durable - def handler(_event: Any, _ctx: DurableContext) -> Any: + def handler(_event: object, _ctx: DurableContext) -> object: return None with pytest.raises(DurableRuntimeMissingError) as raised: _ = handler({}, None) - message = str(raised.value) - # Volcano installs the runtime when it builds a durable function, so the - # fix is a deploy rather than an install. A function's requirements.txt - # never names the runtime, and the error must not send a reader to add it. - assert "deploy this one that way" in message - assert "kind: durable" in message - assert "does not run locally" in message - assert "requirements.txt" not in message - assert "aws-durable-execution-sdk-python" not in message - - # The extra is still the answer for one case, and only that one. - assert "in your own tests, install `volcano-sdk-python[durable]`" in message + assert str(raised.value) == ( + "Durable execution is not available here. Volcano provides the " + "durable runtime when it builds a function deployed as durable, so " + "deploy this one that way (`volcano cloud durable deploy`, or " + "`kind: durable` in volcano-config.yaml). Durable execution is a " + "cloud capability and does not run locally; to exercise a handler " + "in your own tests, install `volcano-sdk-python[durable]`." + ) + assert isinstance(raised.value.__cause__, ImportError) + + +def test_runtime_missing_error_preserves_an_explicit_cause() -> None: + cause = ImportError("missing runtime") + + assert DurableRuntimeMissingError(cause).__cause__ is cause diff --git a/uv.lock b/uv.lock index d2c094a2..2212ea41 100644 --- a/uv.lock +++ b/uv.lock @@ -2217,7 +2217,7 @@ requires-dist = [ { name = "aws-durable-execution-sdk-python", marker = "extra == 'durable'", specifier = ">=2.0.0,<3.0.0" }, { name = "centrifuge-python", specifier = ">=0.6.0,<0.7.0" }, { name = "httpx", specifier = ">=0.28.1,<0.29.0" }, - { name = "typing-extensions", specifier = ">=4.10.0" }, + { name = "typing-extensions", specifier = ">=4.12.2" }, ] provides-extras = ["durable"] From ba551c83ce999e2c6cae6e2714eb88343f8a7a6c Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 05:55:20 -0400 Subject: [PATCH 2/8] test(python): fail durable adapter defects before scheduling --- tests/unit/fixtures/durable_context.py | 17 ++++++--- tests/unit/fixtures/invalid_callbacks.py | 4 ++ tests/unit/test_durable_authoring.py | 48 ++++++++++++++++++++++-- 3 files changed, 60 insertions(+), 9 deletions(-) diff --git a/tests/unit/fixtures/durable_context.py b/tests/unit/fixtures/durable_context.py index 7ebe9ee5..ac022a0f 100644 --- a/tests/unit/fixtures/durable_context.py +++ b/tests/unit/fixtures/durable_context.py @@ -1,4 +1,4 @@ -"""A typed runtime context that records parallel calls without scheduling them.""" +"""A typed runtime context that records batch calls without scheduling them.""" from __future__ import annotations @@ -18,7 +18,7 @@ T = TypeVar("T") U = TypeVar("U") -_UNEXPECTED_OPERATION = "unexpected runtime operation in parallel test" +_UNEXPECTED_OPERATION = "unexpected runtime operation in batch test" class EmptyBatch(Generic[T]): @@ -45,7 +45,7 @@ def throw_if_error(self) -> None: class RecordingContext: - """Implement the runtime context and capture its parallel call.""" + """Implement the runtime context and capture its batch calls.""" logger: DurableLogger = logging.getLogger(__name__) @@ -53,6 +53,8 @@ def __init__(self) -> None: self.branches: list[object] | None = None self.name: str | None = None self.config: object = None + self.map_items: list[object] | None = None + self.map_result: object = None def step( self, @@ -89,8 +91,13 @@ def map( name: str | None, config: object, ) -> _RuntimeBatch[T]: - _ = (items, func, name, config) - raise AssertionError(_UNEXPECTED_OPERATION) + if not items: + raise AssertionError(_UNEXPECTED_OPERATION) + self.name = name + self.config = config + self.map_items = list(items) + self.map_result = func(self, items[0], 7, items) + return EmptyBatch[T]() def parallel( self, diff --git a/tests/unit/fixtures/invalid_callbacks.py b/tests/unit/fixtures/invalid_callbacks.py index f6b67fae..d2bf3080 100644 --- a/tests/unit/fixtures/invalid_callbacks.py +++ b/tests/unit/fixtures/invalid_callbacks.py @@ -23,6 +23,10 @@ def register_non_callable_branch(context: DurableContext) -> None: _ = context.parallel(["not a branch"]) # type: ignore[list-item] +def register_non_callable_map(context: DurableContext) -> None: + _ = context.map([1], None) # type: ignore[arg-type] + + def run_non_callable_operation(context: DurableContext, operation: str) -> object: if operation == "step": return context.step("named", "not a function") # type: ignore[arg-type] diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index 4043354f..c9762b2a 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -31,6 +31,7 @@ from fixtures.invalid_callbacks import ( decorate_non_callable, register_non_callable_branch, + register_non_callable_map, run_non_callable_operation, use_non_callable_retry, ) @@ -77,6 +78,10 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: return isinstance(value, list) +def _is_map_config(value: object) -> TypeGuard[MapConfig[object]]: + return isinstance(value, MapConfig) + + def _ready(state: object) -> bool: if not _is_mapping(state): msg = "expected a state mapping" @@ -116,6 +121,36 @@ def test_map_options_forward_both_batch_limits() -> None: assert config.completion_config.min_successful == 1 +def test_map_forwards_items_callback_index_name_and_batch_limits() -> None: + runtime = RecordingContext() + context = DurableContext(runtime, durable_authoring._Engine()) + observed: list[tuple[int, int]] = [] + + def run(item: int, _child: DurableContext, index: int) -> int: + observed.append((item, index)) + return item * 10 + + result = context.map( + [2, 3], run, "batch", BatchOptions(concurrency=2, min_succeeded=1) + ) + + assert result.completed == 0 + assert runtime.map_items == [2, 3] + assert runtime.map_result == 20 + assert observed == [(2, 7)] + assert runtime.name == "batch" + assert _is_map_config(runtime.config) + assert runtime.config.max_concurrency == 2 + assert runtime.config.completion_config.min_successful == 1 + + +def test_map_requires_a_callable() -> None: + context = DurableContext(RecordingContext(), durable_authoring._Engine()) + + with pytest.raises(TypeError, match=r"map\(\) requires a function to run"): + register_non_callable_map(context) + + def test_parallel_forwards_branches_name_and_batch_limits() -> None: runtime = RecordingContext() context = DurableContext(runtime, durable_authoring._Engine()) @@ -152,6 +187,12 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne # Fail an unavailable or malformed runtime before entering the scheduler: # it otherwise waits for a result that the handler cannot produce. assert_runtime_surface() + # An invalid no-retry decision can leave the scheduler waiting forever. + no_retry = durable_authoring._Engine.load()._never_retry()( + RuntimeError("preflight"), 1 + ) + assert no_retry.should_retry is False + assert no_retry.delay.to_seconds() == 0 runner = DurableFunctionTestRunner(handler) try: yield runner @@ -827,11 +868,10 @@ def handler(_event: object, ctx: DurableContext) -> object: def test_parallel_refuses_a_branch_that_is_not_callable() -> None: - @durable - def handler(_event: object, ctx: DurableContext) -> None: - register_non_callable_branch(ctx) + context = DurableContext(RecordingContext(), durable_authoring._Engine()) - assert "a parallel branch is a callable" in failing_handler(handler) + with pytest.raises(TypeError, match="a parallel branch is a callable"): + register_non_callable_branch(context) def test_a_step_refuses_an_unusable_retry() -> None: From e27bd092b12b88d8c84774a234cb2e466df3707f Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 06:31:58 -0400 Subject: [PATCH 3/8] test(python): verify durable batch and wait boundaries --- maintainers/quality-policy.lock.json | 1 + pyproject.toml | 1 + scripts/check_quality_policy.py | 2 +- src/volcano_sdk/durable_authoring.py | 12 +- tests/unit/fixtures/durable_context.py | 53 +++++++- tests/unit/fixtures/invalid_callbacks.py | 7 +- tests/unit/test_durable_authoring.py | 154 ++++++++++++++++++++++- 7 files changed, 220 insertions(+), 10 deletions(-) diff --git a/maintainers/quality-policy.lock.json b/maintainers/quality-policy.lock.json index 59e90f6b..54d94fc7 100644 --- a/maintainers/quality-policy.lock.json +++ b/maintainers/quality-policy.lock.json @@ -303,6 +303,7 @@ "src/volcano_sdk/auth.py", "src/volcano_sdk/storage.py", "src/volcano_sdk/functions.py", + "src/volcano_sdk/durable_authoring.py", "src/volcano_sdk/database.py", "src/volcano_sdk/_function_resolution.py", "src/volcano_sdk/_transport.py", diff --git a/pyproject.toml b/pyproject.toml index 8515c15f..d012d250 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -462,6 +462,7 @@ type_check_command = [ "src/volcano_sdk/auth.py", "src/volcano_sdk/storage.py", "src/volcano_sdk/functions.py", + "src/volcano_sdk/durable_authoring.py", "src/volcano_sdk/database.py", "src/volcano_sdk/_function_resolution.py", "src/volcano_sdk/_transport.py", diff --git a/scripts/check_quality_policy.py b/scripts/check_quality_policy.py index b11cb503..784e2d51 100644 --- a/scripts/check_quality_policy.py +++ b/scripts/check_quality_policy.py @@ -18,7 +18,7 @@ from collections.abc import Iterable GENERATED = "src/volcano_sdk/_generated" -LOCK_SHA256 = "4a6e8f9b39cf3a86ae00517ee5cbe957375009bffab8834be4ea20107bad5817" +LOCK_SHA256 = "489b1a41632f67d528dbb10267bb8921c38c78726c430d990b4ee99ebdd9236a" TYPE_FIXTURES = { "tests/typing/contract_steps.py", "tests/typing/durable_callbacks.py", diff --git a/src/volcano_sdk/durable_authoring.py b/src/volcano_sdk/durable_authoring.py index b67f47df..55eb7da9 100644 --- a/src/volcano_sdk/durable_authoring.py +++ b/src/volcano_sdk/durable_authoring.py @@ -638,8 +638,8 @@ def __init__(self, batch: _RuntimeBatch[T]) -> None: BatchItem( index=item.index, status=str(getattr(item.status, "value", item.status)).lower(), - result=getattr(item, "result", None), - error=_batch_failure(getattr(item, "error", None)), + result=item.result, + error=_batch_failure(item.error), ) for item in _batch_items(batch) ) @@ -1178,7 +1178,7 @@ def _parse_duration(text: str, field_name: str) -> int: def _scan(text: str, start: int, accept: Callable[[str], bool]) -> int: - at = start - while at < len(text) and accept(text[at]): - at += 1 - return at + for at in range(start, len(text)): + if not accept(text[at]): + return at + return len(text) diff --git a/tests/unit/fixtures/durable_context.py b/tests/unit/fixtures/durable_context.py index ac022a0f..17e180e5 100644 --- a/tests/unit/fixtures/durable_context.py +++ b/tests/unit/fixtures/durable_context.py @@ -3,8 +3,11 @@ from __future__ import annotations import logging +from dataclasses import dataclass from typing import TYPE_CHECKING, Generic, TypeVar +from typing_extensions import override + if TYPE_CHECKING: from collections.abc import Callable @@ -44,6 +47,52 @@ def throw_if_error(self) -> None: return +@dataclass(frozen=True) +class RecordedFailure: + """Error fields preserved by the public batch result.""" + + message: str + type: str + data: str + + +@dataclass +class RecordedItem: + """One settled item in a recorded batch.""" + + index: int + status: object + result: str | None + error: object + + +class RecordedBatch(EmptyBatch[str]): + """A settled success and failure with plain statuses and error details.""" + + success_count = 1 + failure_count = 1 + completion_reason: object = "FINISHED" + + def __init__(self, error: object) -> None: + self._error = error + + @override + def succeeded(self) -> list[_RuntimeBatchItem[str]]: + return [RecordedItem(1, "SUCCEEDED", "done", None)] + + @override + def failed(self) -> list[_RuntimeBatchItem[str]]: + return [RecordedItem(0, "FAILED", None, self._error)] + + @override + def get_results(self) -> list[str]: + return ["done"] + + @override + def get_errors(self) -> list[object]: + return [self._error] + + class RecordingContext: """Implement the runtime context and capture its batch calls.""" @@ -81,7 +130,9 @@ def wait_for_condition( config: object, name: str | None, ) -> T: - _ = (func, config, name) + _ = func + self.config = config + self.name = name raise AssertionError(_UNEXPECTED_OPERATION) def map( diff --git a/tests/unit/fixtures/invalid_callbacks.py b/tests/unit/fixtures/invalid_callbacks.py index d2bf3080..d2fc60eb 100644 --- a/tests/unit/fixtures/invalid_callbacks.py +++ b/tests/unit/fixtures/invalid_callbacks.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING -from volcano_sdk.durable_authoring import durable +from volcano_sdk.durable_authoring import WaitUntilOptions, durable if TYPE_CHECKING: from volcano_sdk.auth import Auth @@ -27,6 +27,11 @@ def register_non_callable_map(context: DurableContext) -> None: _ = context.map([1], None) # type: ignore[arg-type] +def register_non_callable_wait(context: DurableContext) -> None: + options = WaitUntilOptions(until=lambda state: state, initial_state=False) + _ = context.wait_until(None, options) # type: ignore[arg-type] + + def run_non_callable_operation(context: DurableContext, operation: str) -> object: if operation == "step": return context.step("named", "not a function") # type: ignore[arg-type] diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index c9762b2a..a55fefc3 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -25,13 +25,24 @@ ParallelBranch as EngineParallelBranch, ) from aws_durable_execution_sdk_python.retries import RetryDecision +from aws_durable_execution_sdk_python.waits import ( + WaitForConditionConfig, + WaitForConditionDecision, + WaitStrategyConfig, + create_wait_strategy, +) from aws_durable_execution_sdk_python_testing import DurableFunctionTestRunner -from fixtures.durable_context import RecordingContext +from fixtures.durable_context import ( + RecordedBatch, + RecordedFailure, + RecordingContext, +) from fixtures.durable_engine import assert_runtime_surface from fixtures.invalid_callbacks import ( decorate_non_callable, register_non_callable_branch, register_non_callable_map, + register_non_callable_wait, run_non_callable_operation, use_non_callable_retry, ) @@ -40,7 +51,9 @@ from volcano_sdk import durable_authoring from volcano_sdk.durable_authoring import ( BatchFailure, + BatchItem, BatchOptions, + BatchResult, DurableContext, DurableRuntimeMissingError, FunctionHandler, @@ -82,6 +95,10 @@ def _is_map_config(value: object) -> TypeGuard[MapConfig[object]]: return isinstance(value, MapConfig) +def _is_wait_config(value: object) -> TypeGuard[WaitForConditionConfig[bool]]: + return isinstance(value, WaitForConditionConfig) + + def _ready(state: object) -> bool: if not _is_mapping(state): msg = "expected a state mapping" @@ -111,6 +128,101 @@ def test_retry_false_produces_an_immediate_no_retry_decision() -> None: assert decision.delay.to_seconds() == 0 +def test_custom_retry_receives_the_original_error() -> None: + engine = durable_authoring._Engine() + failure = RuntimeError("failed") + seen: list[Exception] = [] + + def decide(error: Exception, _attempt: int) -> RetryDecision: + seen.append(error) + return RetryDecision(should_retry=False, delay=Duration.from_seconds(0)) + + retry = engine._custom_retry_strategy(decide) + + assert retry(failure, 1).should_retry is False + assert seen == [failure] + assert seen[0] is failure + + +def test_wait_options_forward_predicate_timing_and_attempt_budget( + monkeypatch: pytest.MonkeyPatch, +) -> None: + engine = durable_authoring._Engine() + captured: list[WaitStrategyConfig[bool]] = [] + + def record( + config: WaitStrategyConfig[bool], + ) -> Callable[[bool, int], WaitForConditionDecision]: + captured.append(config) + return create_wait_strategy(config) + + monkeypatch.setattr(engine, "create_wait_strategy", record) + configured = engine.wait_condition_options( + WaitUntilOptions( + until=lambda state: state, + initial_state=False, + interval="3s", + max_interval="8s", + backoff_rate=1.25, + max_attempts=7, + ) + ) + + assert len(captured) == 1 + config = captured[0] + incomplete = False + complete = True + assert config.should_continue_polling(incomplete) is True + assert config.should_continue_polling(complete) is False + assert config.max_attempts == 7 + assert config.initial_delay.to_seconds() == 3 + assert config.max_delay.to_seconds() == 8 + assert config.backoff_rate == pytest.approx(1.25) + assert _is_wait_config(configured) + assert configured.initial_state is False + + +def test_wait_options_name_invalid_timing_fields() -> None: + engine = durable_authoring._Engine() + + with pytest.raises(ValueError, match=r"^interval must be a duration"): + _ = engine.wait_condition_options( + WaitUntilOptions( + until=lambda state: state, initial_state=False, interval="bad" + ) + ) + with pytest.raises(ValueError, match=r"^max_interval must be a duration"): + _ = engine.wait_condition_options( + WaitUntilOptions( + until=lambda state: state, + initial_state=False, + max_interval="bad", + ) + ) + + +def test_wait_until_validates_callback_and_forwards_name() -> None: + runtime = RecordingContext() + context = DurableContext(runtime, durable_authoring._Engine()) + options = WaitUntilOptions(until=lambda state: state, initial_state=False) + + with pytest.raises(TypeError, match=r"wait_until\(\) requires a function"): + register_non_callable_wait(context) + with pytest.raises(AssertionError, match="unexpected runtime operation"): + _ = context.wait_until(lambda state, _scope: state, options, "poll-ready") + assert runtime.name == "poll-ready" + + +def test_wait_accepts_the_maximum_and_names_an_invalid_duration() -> None: + context = DurableContext(RecordingContext(), durable_authoring._Engine()) + maximum = context._wait_duration(31_622_400) + + assert isinstance(maximum, Duration) + assert maximum.to_seconds() == 31_622_400 + with pytest.raises(ValueError, match=r"^wait must be a duration"): + _ = context._wait_duration("bad") + + def test_map_options_forward_both_batch_limits() -> None: config = durable_authoring._Engine().map_options( BatchOptions(concurrency=2, min_succeeded=1) @@ -151,6 +263,32 @@ def test_map_requires_a_callable() -> None: register_non_callable_map(context) +def test_batch_result_preserves_settled_items_and_failure_details() -> None: + failure = BatchFailure("failed", "RemoteError", "trace") + result = BatchResult( + RecordedBatch(RecordedFailure("failed", "RemoteError", "trace")) + ) + + assert result.items == ( + BatchItem(0, "failed", None, failure), + BatchItem(1, "succeeded", "done", None), + ) + assert result.results == ("done",) + assert result.errors == (failure,) + assert result.succeeded == 1 + assert result.failed == 1 + assert result.completed == 2 + assert result.completion_reason == "finished" + + +def test_batch_result_keeps_a_plain_exception_message() -> None: + failure = BatchFailure("boom") + result = BatchResult(RecordedBatch(RuntimeError("boom"))) + + assert result.items[0].error == failure + assert result.errors == (failure,) + + def test_parallel_forwards_branches_name_and_batch_limits() -> None: runtime = RecordingContext() context = DurableContext(runtime, durable_authoring._Engine()) @@ -590,6 +728,20 @@ def test_retry_config_keeps_unset_defaults_and_sets_requested_fields() -> None: assert config.retryable_error_types == [ValueError] +@pytest.mark.parametrize( + ("options", "field"), + [ + (RetryOptions(initial_delay="invalid"), "initial_delay"), + (RetryOptions(max_delay="invalid"), "max_delay"), + ], +) +def test_retry_config_names_an_invalid_duration( + options: RetryOptions, field: str +) -> None: + with pytest.raises(ValueError, match=rf"^{field} must be a duration"): + _ = durable_authoring._Engine()._retry_config(options) + + def test_custom_retry_must_return_a_runtime_decision() -> None: engine = durable_authoring._Engine.load() From 1c749fe34e084f0813f2516845100b36b31ef95c Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 06:46:18 -0400 Subject: [PATCH 4/8] test(python): fail missing wait budgets before scheduling --- tests/unit/test_durable_authoring.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index a55fefc3..fecebace 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -24,6 +24,7 @@ from aws_durable_execution_sdk_python.config import ( ParallelBranch as EngineParallelBranch, ) +from aws_durable_execution_sdk_python.exceptions import WaitForConditionError from aws_durable_execution_sdk_python.retries import RetryDecision from aws_durable_execution_sdk_python.waits import ( WaitForConditionConfig, @@ -808,16 +809,23 @@ def check( def test_wait_until_fails_when_it_runs_out_of_attempts() -> None: + options = WaitUntilOptions( + until=lambda state: state, + initial_state=False, + interval="1s", + max_attempts=2, + ) + configured = durable_authoring._Engine().wait_condition_options(options) + assert _is_wait_config(configured) + not_ready = False + with pytest.raises(WaitForConditionError, match="exhausted 2 attempts"): + _ = configured.wait_strategy(not_ready, 2) + @durable def handler(_event: object, ctx: DurableContext) -> object: return ctx.wait_until( - lambda _state, _scope: {"ready": False}, - WaitUntilOptions( - until=_ready, - initial_state={"ready": False}, - interval="1s", - max_attempts=2, - ), + lambda _state, _scope: False, + options, "never-ready", ) From eafdcd2b2e8a9e22529c79bb36cbc41429106466 Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:20:23 -0400 Subject: [PATCH 5/8] test(python): bound durable duration mutation failures --- src/volcano_sdk/durable_authoring.py | 10 ++++------ tests/unit/test_durable_authoring.py | 13 +++++++++++++ 2 files changed, 17 insertions(+), 6 deletions(-) diff --git a/src/volcano_sdk/durable_authoring.py b/src/volcano_sdk/durable_authoring.py index 55eb7da9..cefdbd4f 100644 --- a/src/volcano_sdk/durable_authoring.py +++ b/src/volcano_sdk/durable_authoring.py @@ -1153,8 +1153,7 @@ def _parse_duration(text: str, field_name: str) -> int: ValueError: The text is empty or contains an invalid number or unit. """ - seconds = 0 - segments = 0 + parts: list[int] = [] at = 0 while at < len(text): if text[at] == " ": @@ -1165,16 +1164,15 @@ def _parse_duration(text: str, field_name: str) -> int: unit = text[number_end:unit_end] if number_end == at or unit not in _DURATION_UNITS: break - seconds += int(text[at:number_end]) * _DURATION_UNITS[unit] - segments += 1 + parts.append(int(text[at:number_end]) * _DURATION_UNITS[unit]) at = unit_end - if segments == 0 or at != len(text): + if not parts or at != len(text): message = ( f"{field_name} must be a duration in whole seconds, such as '30s', " f"'5m', '2h', '1d' or '1m30s' (got {text!r})" ) raise ValueError(message) - return seconds + return sum(parts) def _scan(text: str, start: int, accept: Callable[[str], bool]) -> int: diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index fecebace..035c46a6 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -20,6 +20,7 @@ Duration, MapConfig, ParallelConfig, + StepSemantics, ) from aws_durable_execution_sdk_python.config import ( ParallelBranch as EngineParallelBranch, @@ -326,6 +327,7 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne # Fail an unavailable or malformed runtime before entering the scheduler: # it otherwise waits for a result that the handler cannot produce. assert_runtime_surface() + assert to_seconds({"seconds": 5}, "wait") == 5 # An invalid no-retry decision can leave the scheduler waiting forever. no_retry = durable_authoring._Engine.load()._never_retry()( RuntimeError("preflight"), 1 @@ -776,6 +778,9 @@ def flaky(scope: StepScope) -> int: def test_at_most_once_runs_the_step_once() -> None: + config = durable_authoring._Engine().step_options(retry=None, at_most_once=True) + assert config.step_semantics is StepSemantics.AT_MOST_ONCE_PER_RETRY + @durable def handler(_event: object, ctx: DurableContext) -> object: return ctx.step("charge", lambda _scope: "charged", at_most_once=True) @@ -1105,6 +1110,14 @@ def test_an_unknown_duration_field_names_what_it_accepts() -> None: @pytest.mark.parametrize( ("value", "error_type", "message"), [ + ( + None, + TypeError, + ( + "interval must be a duration string, a whole number of seconds, " + "or a mapping of days, hours, minutes, seconds" + ), + ), ( True, TypeError, From fffa3f4aebc8305fb1c92490a4592ac0c015954b Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 07:45:29 -0400 Subject: [PATCH 6/8] test(python): bound durable wait predicate defects --- tests/unit/test_durable_authoring.py | 14 +++++++++++--- 1 file changed, 11 insertions(+), 3 deletions(-) diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index 035c46a6..b010f30c 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -328,10 +328,18 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne # it otherwise waits for a result that the handler cannot produce. assert_runtime_surface() assert to_seconds({"seconds": 5}, "wait") == 5 - # An invalid no-retry decision can leave the scheduler waiting forever. - no_retry = durable_authoring._Engine.load()._never_retry()( - RuntimeError("preflight"), 1 + assert to_seconds("1s", "wait") == 1 + engine = durable_authoring._Engine.load() + wait_config = engine.wait_condition_options( + WaitUntilOptions(until=lambda state: state, initial_state=False, max_attempts=2) ) + assert _is_wait_config(wait_config) + not_ready = False + ready = True + assert wait_config.wait_strategy(not_ready, 1).should_continue is True + assert wait_config.wait_strategy(ready, 1).should_continue is False + # An invalid no-retry decision can leave the scheduler waiting forever. + no_retry = engine._never_retry()(RuntimeError("preflight"), 1) assert no_retry.should_retry is False assert no_retry.delay.to_seconds() == 0 runner = DurableFunctionTestRunner(handler) From 3e815a57da2c28e807e284c7d3825a33683ae920 Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 08:16:17 -0400 Subject: [PATCH 7/8] test(python): assert durable step defaults and bound duration parsing --- src/volcano_sdk/durable_authoring.py | 9 +++++---- tests/unit/fixtures/durable_context.py | 3 ++- tests/unit/test_durable_authoring.py | 12 ++++++++++++ 3 files changed, 19 insertions(+), 5 deletions(-) diff --git a/src/volcano_sdk/durable_authoring.py b/src/volcano_sdk/durable_authoring.py index cefdbd4f..84f7f354 100644 --- a/src/volcano_sdk/durable_authoring.py +++ b/src/volcano_sdk/durable_authoring.py @@ -1155,10 +1155,11 @@ def _parse_duration(text: str, field_name: str) -> int: """ parts: list[int] = [] at = 0 - while at < len(text): - if text[at] == " ": - at += 1 - continue + # Each valid iteration consumes input, so its length bounds the scan. + for _ in text: + at = _scan(text, at, lambda char: char == " ") + if at == len(text): + break number_end = _scan(text, at, str.isdigit) unit_end = _scan(text, number_end, str.islower) unit = text[number_end:unit_end] diff --git a/tests/unit/fixtures/durable_context.py b/tests/unit/fixtures/durable_context.py index 17e180e5..185a4b27 100644 --- a/tests/unit/fixtures/durable_context.py +++ b/tests/unit/fixtures/durable_context.py @@ -111,7 +111,8 @@ def step( name: str | None, config: object, ) -> T: - _ = (func, name, config) + _ = (func, name) + self.config = config raise AssertionError(_UNEXPECTED_OPERATION) def wait(self, duration: object, name: str | None = None) -> None: diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index b010f30c..c2135793 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -20,6 +20,7 @@ Duration, MapConfig, ParallelConfig, + StepConfig, StepSemantics, ) from aws_durable_execution_sdk_python.config import ( @@ -796,6 +797,17 @@ def handler(_event: object, ctx: DurableContext) -> object: assert run_handler(handler) == "charged" +def test_step_uses_at_least_once_semantics_by_default() -> None: + runtime = RecordingContext() + context = DurableContext(runtime, durable_authoring._Engine()) + + with pytest.raises(AssertionError, match="unexpected runtime operation"): + _ = context.step("default", lambda _scope: None) + + assert isinstance(runtime.config, StepConfig) + assert runtime.config.step_semantics is StepSemantics.AT_LEAST_ONCE_PER_RETRY + + def test_wait_until_polls_until_the_condition_holds() -> None: polls = {"count": 0} From 4e783b28f246ffdea6a99f825443b0ea5ddca592 Mon Sep 17 00:00:00 2001 From: Sean Keever <33592180+swkeever@users.noreply.github.com> Date: Thu, 24 Sep 2026 08:49:39 -0400 Subject: [PATCH 8/8] test(python): preflight durable retry policy before scheduling --- tests/unit/test_durable_authoring.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/tests/unit/test_durable_authoring.py b/tests/unit/test_durable_authoring.py index c2135793..64178228 100644 --- a/tests/unit/test_durable_authoring.py +++ b/tests/unit/test_durable_authoring.py @@ -339,8 +339,14 @@ def local_runner(handler: FunctionHandler) -> Generator[DurableFunctionTestRunne ready = True assert wait_config.wait_strategy(not_ready, 1).should_continue is True assert wait_config.wait_strategy(ready, 1).should_continue is False - # An invalid no-retry decision can leave the scheduler waiting forever. - no_retry = engine._never_retry()(RuntimeError("preflight"), 1) + # An invalid retry policy can leave the scheduler waiting forever. + assert engine.step_options(retry=None, at_most_once=False).retry_strategy is None + assert engine.step_options(retry=True, at_most_once=False).retry_strategy is None + no_retry_strategy = engine.step_options( + retry=False, at_most_once=False + ).retry_strategy + assert no_retry_strategy is not None + no_retry = no_retry_strategy(RuntimeError("preflight"), 1) assert no_retry.should_retry is False assert no_retry.delay.to_seconds() == 0 runner = DurableFunctionTestRunner(handler)