diff --git a/packages/aws-durable-execution-sdk-python-testing/README.md b/packages/aws-durable-execution-sdk-python-testing/README.md index d6e4dc74c..58b21e793 100644 --- a/packages/aws-durable-execution-sdk-python-testing/README.md +++ b/packages/aws-durable-execution-sdk-python-testing/README.md @@ -12,6 +12,7 @@ - [Installation](#installation) - [Quick Start](#quick-start) +- [Testing functions that invoke other functions](#testing-functions-that-invoke-other-functions) - [Architecture](#architecture) - [Documentation](#documentation) - [Developer Guide](#developers) @@ -111,6 +112,55 @@ def test_my_durable_functions(): three_result: StepOperation = result.get_step("three") assert three_result.result == '"5 6"' ``` +## Testing functions that invoke other functions + +A durable function can call another function with `context.invoke`. The +runner needs to know each target's kind: a durable target runs as a +child durable execution with its own history, while a plain (non-durable) +target is one invocation whose return value is the result. An unknown +target fails the operation with `ResourceNotFoundException`. + +### In-process runner + +Register each target with the runner. `register_durable_function` marks a +durable target, and takes its execution timeout and retention; +`register_function` marks a plain one. The runner has one handler per +registered name, not versions: an invoke of `child:prod` runs the handler +registered as `child:prod` if there is one, otherwise the one registered +as `child`. + +```python +from aws_durable_execution_sdk_python.context import DurableContext +from aws_durable_execution_sdk_python.execution import durable_execution +from aws_durable_execution_sdk_python_testing.runner import DurableFunctionTestRunner + +@durable_execution +def process_payment(event: dict, context: DurableContext) -> dict: + return {"charged": event["amount"]} + +def lookup_price(event: dict, context) -> dict: + return {"price": 25} + +@durable_execution +def place_order(event: dict, context: DurableContext) -> dict: + price = context.invoke("lookup-price", {"sku": event["sku"]}, name="price") + payment = context.invoke( + "process-payment", {"amount": price["price"]}, name="payment" + ) + return {"sku": event["sku"], "charged": payment["charged"]} + +def test_place_order_invokes_both_functions(): + with DurableFunctionTestRunner(handler=place_order) as runner: + runner.register_durable_function( + "process-payment", process_payment, execution_timeout=60 + ) + runner.register_function("lookup-price", lookup_price) + result = runner.run(input='{"sku": "book-1"}') + + assert result.result == '{"sku": "book-1", "charged": 25}' + assert result.get_invoke("payment").status.value == "SUCCEEDED" +``` + ## Architecture See [docs/architecture.md](docs/architecture.md) for framework diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/effects.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/effects.py index 69d642e57..4dd6dc282 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/effects.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/effects.py @@ -47,4 +47,15 @@ class CallbackCreated: callback_token: CallbackToken -CheckpointEffect = Completed | Failed | CallbackCreated +@dataclass(frozen=True) +class ChainedInvokeStarted: + """A chained invoke was accepted and its target must be dispatched.""" + + execution_arn: str + operation_id: str + function_name: str + tenant_id: str | None + payload: str | None + + +CheckpointEffect = Completed | Failed | CallbackCreated | ChainedInvokeStarted diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processor.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processor.py index d9360e0a5..42ea9333b 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processor.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processor.py @@ -29,7 +29,12 @@ if TYPE_CHECKING: - from aws_durable_execution_sdk_python.lambda_service import OperationUpdate + from collections.abc import Callable + + from aws_durable_execution_sdk_python.lambda_service import ( + ErrorObject, + OperationUpdate, + ) from aws_durable_execution_sdk_python_testing.clock import Clock from aws_durable_execution_sdk_python_testing.execution import Execution @@ -68,6 +73,18 @@ def add_execution_observer(self, observer: ExecutionObserver) -> None: """Add observer for execution events.""" self._observers.append(observer) + def set_chained_invoke_preflight( + self, preflight: Callable[[str], ErrorObject | None] + ) -> None: + """Resolve chained-invoke targets at checkpoint time with ``preflight``. + + A target it fails comes back FAILED in the checkpoint response + instead of being dispatched (see ``ChainedInvokeProcessor``). + """ + self._dispatcher = CheckpointRequestDispatcher( + chained_invoke_preflight=preflight + ) + def process_checkpoint( self, checkpoint_token: str, diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/base.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/base.py index 52eb64c24..d2eaa4ec4 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/base.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/base.py @@ -44,6 +44,19 @@ def process( """ raise NotImplementedError + def stored_update( + self, + update: OperationUpdate, + current_op: Operation | None, # noqa: ARG002 + updated_op: Operation, # noqa: ARG002 + ) -> OperationUpdate: + """The update to keep in the execution's record for history. + + Most processors keep the update as sent. A processor overrides + this when the service keeps less than it was sent. + """ + return update + def _get_start_time( self, current_operation: Operation | None, now: datetime.datetime ) -> datetime.datetime | None: diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/invoke.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/invoke.py new file mode 100644 index 000000000..1e6d75649 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/processors/invoke.py @@ -0,0 +1,139 @@ +"""Chained invoke operation processor for handling CHAINED_INVOKE operation updates.""" + +from __future__ import annotations + +import datetime +from dataclasses import replace +from typing import TYPE_CHECKING + +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeDetails, + ErrorObject, + Operation, + OperationAction, + OperationStatus, + OperationUpdate, +) + +from aws_durable_execution_sdk_python_testing.checkpoint.processors.base import ( + OperationProcessor, +) +from aws_durable_execution_sdk_python_testing.exceptions import ( + InvalidParameterValueException, +) + +if TYPE_CHECKING: + from collections.abc import Callable + + from aws_durable_execution_sdk_python_testing.observer import ExecutionNotifier + + +class ChainedInvokeProcessor(OperationProcessor): + """Processes CHAINED_INVOKE operation updates. + + The service resolves the target before it schedules anything. Its + checkpoint response then carries the operation STARTED, because the + scheduling event is written as part of completing the checkpoint, + or FAILED when the target could not be resolved, so the handler + sees such a failure in the response to its own checkpoint and + raises without suspending. This processor does the same with + ``preflight``: a START update whose target fails the preflight + comes back FAILED with that error; otherwise it comes back STARTED + and raises a :class:`ChainedInvokeStarted` effect for the caller to + dispatch the target once the checkpoint write completes. The + operation holds STARTED until its terminal transition, so a handler + observing it sees exactly one status change; the target's terminal + state completes it. Completion is never a handler checkpoint, so + START is the only accepted action. + """ + + def __init__(self, preflight: Callable[[str], ErrorObject | None] | None = None): + """``preflight`` resolves a target at checkpoint time; without it + every target is dispatched and any failure arrives later.""" + self._preflight = preflight + + def process( + self, + update: OperationUpdate, + current_op: Operation | None, + notifier: ExecutionNotifier, + execution_arn: str, + now: datetime.datetime, + ) -> Operation: + """Process CHAINED_INVOKE operation update.""" + match update.action: + case OperationAction.START: + options = update.chained_invoke_options + if options is None: + msg_options_required: str = ( + "Update for CHAINED_INVOKE operation requires " + "ChainedInvokeOptions." + ) + raise InvalidParameterValueException(msg_options_required) + + start_timestamp: datetime.datetime | None = self._get_start_time( + current_op, now + ) + error: ErrorObject | None = ( + self._preflight(options.function_name) + if self._preflight is not None + else None + ) + if error is not None: + return Operation( + operation_id=update.operation_id, + parent_id=update.parent_id, + name=update.name, + start_timestamp=start_timestamp, + end_timestamp=now, + operation_type=update.operation_type, + status=OperationStatus.FAILED, + sub_type=update.sub_type, + chained_invoke_details=ChainedInvokeDetails( + result=None, error=error + ), + ) + + operation: Operation = Operation( + operation_id=update.operation_id, + parent_id=update.parent_id, + name=update.name, + start_timestamp=start_timestamp, + end_timestamp=None, + operation_type=update.operation_type, + status=OperationStatus.STARTED, + sub_type=update.sub_type, + chained_invoke_details=ChainedInvokeDetails( + result=None, error=None + ), + ) + + notifier.notify_chained_invoke_started( + execution_arn=execution_arn, + operation_id=update.operation_id, + function_name=options.function_name, + tenant_id=options.tenant_id, + payload=update.payload, + ) + return operation + case _: + msg_invalid_action: str = "Invalid action for CHAINED_INVOKE operation." + raise InvalidParameterValueException(msg_invalid_action) + + def stored_update( + self, + update: OperationUpdate, + current_op: Operation | None, + updated_op: Operation, + ) -> OperationUpdate: + """Keep no input for a chained invoke that failed before starting. + + The service stores no input payload when the target cannot be + resolved: it records a payload size of zero and an empty payload + location, and its history shows the function name and the error + without the input. So the runner stores the update without its + payload, which also keeps it out of the operation's size. + """ + if current_op is None and updated_op.status is OperationStatus.FAILED: + return replace(update, payload=None) + return update diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/transformer.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/transformer.py index 62971fd64..ff219a17a 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/transformer.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/transformer.py @@ -23,6 +23,9 @@ from aws_durable_execution_sdk_python_testing.checkpoint.processors.execution import ( ExecutionProcessor, ) +from aws_durable_execution_sdk_python_testing.checkpoint.processors.invoke import ( + ChainedInvokeProcessor, +) from aws_durable_execution_sdk_python_testing.checkpoint.processors.step import ( StepProcessor, ) @@ -40,6 +43,7 @@ from datetime import datetime from aws_durable_execution_sdk_python.lambda_service import ( + ErrorObject, OperationUpdate, ) @@ -69,13 +73,28 @@ class CheckpointRequestDispatcher: OperationType.CONTEXT: ContextProcessor(), OperationType.CALLBACK: CallbackProcessor(), OperationType.EXECUTION: ExecutionProcessor(), + OperationType.CHAINED_INVOKE: ChainedInvokeProcessor(), } def __init__( self, processors: MutableMapping[OperationType, OperationProcessor] | None = None, + *, + chained_invoke_preflight: Callable[[str], ErrorObject | None] | None = None, ): - self.processors = processors if processors else self._DEFAULT_PROCESSORS + """``chained_invoke_preflight`` resolves a chained-invoke target at + checkpoint time (see :class:`ChainedInvokeProcessor`); it replaces + the default CHAINED_INVOKE processor and is ignored when explicit + ``processors`` are given.""" + if processors: + self.processors = processors + elif chained_invoke_preflight is not None: + self.processors = dict(self._DEFAULT_PROCESSORS) + self.processors[OperationType.CHAINED_INVOKE] = ChainedInvokeProcessor( + chained_invoke_preflight + ) + else: + self.processors = self._DEFAULT_PROCESSORS def apply_updates( self, @@ -112,6 +131,7 @@ def apply_updates( """ collector = ExecutionNotifier() op_map = {op.operation_id: op for op in execution.operations} + stored_updates: list[OperationUpdate] = [] for update in updates: processor = self.processors.get(update.operation_type) @@ -128,8 +148,11 @@ def apply_updates( now=now, ) if updated_op is None: + stored_updates.append(update) continue + stored = processor.stored_update(update, current_op, updated_op) + stored_updates.append(stored) if update.operation_id in op_map: for i, op in enumerate(execution.operations): # pragma: no branch if op.operation_id == update.operation_id: @@ -140,12 +163,12 @@ def apply_updates( op_map[update.operation_id] = updated_op execution.operation_size_bytes[update.operation_id] = ( - _estimate_payload_size(update) + _estimate_payload_size(stored) ) touch(update.operation_id) - execution.updates.extend(updates) - execution.update_timestamps.extend(now for _ in updates) + execution.updates.extend(stored_updates) + execution.update_timestamps.extend(now for _ in stored_updates) return collector.effects diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/checkpoint.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/checkpoint.py index ab6d80bfe..f5ae3acff 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/checkpoint.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/checkpoint.py @@ -35,6 +35,12 @@ WaitOperationValidator, VALID_ACTIONS_FOR_WAIT, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + CHAINED_INVOKE_INPUT_TOO_LARGE_MESSAGE, + CHILD_EXECUTION_OUTPUT_TOO_LARGE_MESSAGE, + MAX_CHAINED_INVOKE_PAYLOAD_BYTES, + payload_size_bytes, +) from aws_durable_execution_sdk_python_testing.exceptions import ( InvalidParameterValueException, ) @@ -99,20 +105,40 @@ def _validate_operation_update( ) -> None: """Validate a single operation update.""" CheckpointValidator._validate_inconsistent_operation_metadata(update, execution) - CheckpointValidator._validate_payload_sizes(update) + CheckpointValidator._validate_payload_sizes(update, execution) CheckpointValidator._validate_valid_action_for_type( update.operation_type, update.action ) CheckpointValidator._validate_operation_status_transition(update, execution) @staticmethod - def _validate_payload_sizes(update: OperationUpdate) -> None: - """Validate that operation payload sizes are not too large.""" + def _validate_payload_sizes(update: OperationUpdate, execution: Execution) -> None: + """Validate that operation payload sizes are not too large. + + Besides the error-object bound, two chained-invoke bounds apply: + the input a parent sends to a target, and the result a child + execution returns to its parent, are each capped at 1 MiB. + """ if update.error is not None: payload = json.dumps(update.error.to_dict()) if len(payload) > MAX_ERROR_PAYLOAD_SIZE_BYTES: msg: str = f"Error object size must be less than {MAX_ERROR_PAYLOAD_SIZE_BYTES} bytes." raise InvalidParameterValueException(msg) + if update.payload is None: + return + if ( + update.operation_type == OperationType.CHAINED_INVOKE + and payload_size_bytes(update.payload) > MAX_CHAINED_INVOKE_PAYLOAD_BYTES + ): + raise InvalidParameterValueException(CHAINED_INVOKE_INPUT_TOO_LARGE_MESSAGE) + if ( + update.operation_type == OperationType.EXECUTION + and execution.parent_execution_arn is not None + and payload_size_bytes(update.payload) > MAX_CHAINED_INVOKE_PAYLOAD_BYTES + ): + raise InvalidParameterValueException( + CHILD_EXECUTION_OUTPUT_TOO_LARGE_MESSAGE + ) @staticmethod def _validate_valid_action_for_type( @@ -150,7 +176,9 @@ def _validate_operation_status_transition( case OperationType.CALLBACK: CallbackOperationValidator.validate(current_state, update) case OperationType.CHAINED_INVOKE: - ChainedInvokeOperationValidator.validate(current_state, update) + ChainedInvokeOperationValidator.validate( + current_state, update, execution + ) case OperationType.EXECUTION: ExecutionOperationValidator.validate(update) case _: # pragma: no cover diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/operations/invoke.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/operations/invoke.py index 1c2871284..94444b27f 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/operations/invoke.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/checkpoint/validators/operations/invoke.py @@ -2,55 +2,116 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from aws_durable_execution_sdk_python.lambda_service import ( Operation, OperationAction, - OperationStatus, OperationUpdate, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + CHAINED_INVOKE_DIFFERENT_ACCOUNT_MESSAGE, + CHAINED_INVOKE_DIFFERENT_REGION_MESSAGE, + INVALID_FUNCTION_ARN_MESSAGE_FORMAT, + INVALID_TENANT_ID_MESSAGE, + FunctionTarget, + is_valid_tenant_id, + parse_function_target, +) from aws_durable_execution_sdk_python_testing.exceptions import ( InvalidParameterValueException, ) +if TYPE_CHECKING: + from aws_durable_execution_sdk_python_testing.execution import Execution + VALID_ACTIONS_FOR_INVOKE = frozenset( [ OperationAction.START, - OperationAction.CANCEL, ] ) class ChainedInvokeOperationValidator: - """Validates INVOKE operation transitions.""" + """Validates INVOKE operation transitions. - _ALLOWED_STATUS_TO_CANCEL = frozenset( - [ - OperationStatus.STARTED, - ] - ) + START is the only accepted action: a chained invoke completes through + the invoked function's terminal state, not through a handler + checkpoint. START also validates the options, as the service does: + the target must be a well-formed Lambda function name or ARN in the + parent's account and region, and a TenantId must satisfy the API + constraint. + """ @staticmethod - def validate(current_state: Operation | None, update: OperationUpdate) -> None: + def validate( + current_state: Operation | None, + update: OperationUpdate, + execution: Execution, + ) -> None: """Validate INVOKE operation update.""" match update.action: case OperationAction.START: if current_state is not None: msg_invoke_exists: str = ( - "Cannot start an INVOKE that already exist." + "Cannot start a CHAINED_INVOKE operation that already exists." ) raise InvalidParameterValueException(msg_invoke_exists) - case OperationAction.CANCEL: - if ( - current_state is None - or current_state.status - not in ChainedInvokeOperationValidator._ALLOWED_STATUS_TO_CANCEL + if update.chained_invoke_options is None: + msg_options_required: str = ( + "Update for CHAINED_INVOKE operation requires " + "ChainedInvokeOptions." + ) + raise InvalidParameterValueException(msg_options_required) + ChainedInvokeOperationValidator._validate_target( + update.chained_invoke_options.function_name, execution + ) + # The API model types TenantId as a string. A value of + # another type is a validation error, not a server error. + tenant_id: str | None = update.chained_invoke_options.tenant_id + if tenant_id is not None and ( + not isinstance(tenant_id, str) or not is_valid_tenant_id(tenant_id) ): - msg_invoke_cancel: str = "Cannot cancel an INVOKE that does not exist or has already completed." - raise InvalidParameterValueException(msg_invoke_cancel) + raise InvalidParameterValueException(INVALID_TENANT_ID_MESSAGE) case _: - msg_invoke_invalid: str = "Invalid INVOKE action." + msg_invoke_invalid: str = "Invalid action for CHAINED_INVOKE operation." raise InvalidParameterValueException(msg_invoke_invalid) + + @staticmethod + def _validate_target(function_name: str, execution: Execution) -> None: + """Reject a malformed target or one in another account or region. + + The region check applies only when the execution recorded a + region; an execution built without one skips it. + """ + # The API model requires FunctionName as a string. A missing or + # non-string value is a validation error, not a server error. + if not isinstance(function_name, str): + raise InvalidParameterValueException( + INVALID_FUNCTION_ARN_MESSAGE_FORMAT.format(function_name) + ) + try: + target: FunctionTarget = parse_function_target(function_name) + except ValueError as err: + raise InvalidParameterValueException( + INVALID_FUNCTION_ARN_MESSAGE_FORMAT.format(function_name) + ) from err + if ( + target.account_id is not None + and target.account_id != execution.start_input.account_id + ): + raise InvalidParameterValueException( + CHAINED_INVOKE_DIFFERENT_ACCOUNT_MESSAGE + ) + if ( + target.region is not None + and execution.region is not None + and target.region != execution.region + ): + raise InvalidParameterValueException( + CHAINED_INVOKE_DIFFERENT_REGION_MESSAGE + ) diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/child_dispatcher.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/child_dispatcher.py new file mode 100644 index 000000000..39ce8c9ba --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/child_dispatcher.py @@ -0,0 +1,828 @@ +"""Dispatch of chained-invoke targets. + +A CHAINED_INVOKE operation names a function to invoke. The in-process +runner resolves the name from registered handlers; the web runner +resolves it from a function configuration file and invokes at a +Lambda-compatible endpoint. A :class:`ChildDispatcher` owns that +difference and reports one of three results: + +* ``StartChild``: a durable target; the executor creates, links, and + launches a child durable execution; +* ``RunInvocation``: a non-durable target; one invocation whose return + value is the outcome; +* ``KnownOutcome``: the target could not be dispatched. + +A child execution completes the parent's operation through the +executor's terminal-transition hook; no dispatcher waits for it. +""" + +from __future__ import annotations + +import json +import logging +import pathlib +import re +import threading +import uuid +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Final, Protocol + +from botocore.exceptions import ClientError, ReadTimeoutError # type: ignore + +from aws_durable_execution_sdk_python.lambda_service import ErrorObject + +from aws_durable_execution_sdk_python_testing.model import ( + StartDurableExecutionInput, +) + +if TYPE_CHECKING: + from collections.abc import Callable, Mapping + from datetime import datetime + + from aws_durable_execution_sdk_python_testing.clock import Clock + +logger = logging.getLogger(__name__) + +# Default execution settings for a chained child execution when the +# registration does not specify them. +DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS = 120 +DEFAULT_CHILD_RETENTION_PERIOD_DAYS = 1 + +# Error surfaced on a CHAINED_INVOKE operation that timed out: a durable +# child whose execution timed out, or a plain target whose single +# invocation ran past the invocation timeout. The service reports this +# fixed type; a stopped or failed child instead surfaces its own error +# object, and a target that could not be invoked surfaces the Lambda +# API error code. +CHAINED_INVOKE_TIMEOUT_ERROR_TYPE: Final[str] = "ChainedInvoke.Timeout" +# Error type Lambda reports for a function that ran past its timeout. +FUNCTION_TIMEOUT_ERROR_TYPE: Final[str] = "Sandbox.Timedout" + +# Lambda API error code for an Invoke of a function that does not exist. +FUNCTION_NOT_FOUND_ERROR_TYPE: Final[str] = "ResourceNotFoundException" + +# Chained-invoke payload limits, as the service enforces them: the +# input a parent sends, the result a durable child returns, and the +# result a plain target returns are each capped at 1 MiB. +MAX_CHAINED_INVOKE_PAYLOAD_BYTES: Final[int] = 1_048_576 +CHAINED_INVOKE_INPUT_TOO_LARGE_MESSAGE: Final[str] = ( + "CHAINED_INVOKE input payload size must be less than or equal to " + f"{MAX_CHAINED_INVOKE_PAYLOAD_BYTES} bytes." +) +CHAINED_INVOKE_OUTPUT_TOO_LARGE_MESSAGE: Final[str] = ( + "CHAINED_INVOKE output payload size must be less than or equal to " + f"{MAX_CHAINED_INVOKE_PAYLOAD_BYTES} bytes." +) +CHILD_EXECUTION_OUTPUT_TOO_LARGE_MESSAGE: Final[str] = ( + "Execution output payload size must be less than or equal to " + f"{MAX_CHAINED_INVOKE_PAYLOAD_BYTES} bytes." +) +CHAINED_INVOKE_UTF8_DECODING_MESSAGE: Final[str] = ( + "CHAINED_INVOKE response could not be decoded as UTF-8." +) + +# A chained-invoke target is a Lambda function name in any form Invoke +# accepts: bare name, name with qualifier, partial ARN, or full ARN. +# Grammar and length are those of ChainedInvokeOptions.FunctionName. +MAX_FUNCTION_NAME_LENGTH: Final[int] = 256 +_FUNCTION_NAME_PATTERN: Final[re.Pattern[str]] = re.compile( + r"(?Parn:(aws[a-zA-Z-]*)?:lambda:)?" + r"((?P(eusc-)?[a-z]{2}((-gov)|(-iso([a-z]?)))?-[a-z]+-\d{1}):)?" + r"((?P\d{12}):)?" + r"(function:)?" + r"(?P[a-zA-Z0-9-_\.]+)" + r"(:(?P\$LATEST(\.PUBLISHED)?|[a-zA-Z0-9-_]+))?" +) +INVALID_FUNCTION_ARN_MESSAGE_FORMAT: Final[str] = "Invalid function ARN '{}'" + +# ChainedInvokeOptions.TenantId: 1 to 256 characters from this set. +MAX_TENANT_ID_LENGTH: Final[int] = 256 +_TENANT_ID_PATTERN: Final[re.Pattern[str]] = re.compile(r"[a-zA-Z0-9\._:\/=+\-@ ]+") +INVALID_TENANT_ID_MESSAGE: Final[str] = ( + "TenantId must be 1 to 256 characters matching [a-zA-Z0-9._:/=+-@ ]." +) + + +def is_valid_tenant_id(tenant_id: str) -> bool: + """Whether ``tenant_id`` satisfies the ChainedInvokeOptions.TenantId constraint.""" + return ( + 0 < len(tenant_id) <= MAX_TENANT_ID_LENGTH + and _TENANT_ID_PATTERN.fullmatch(tenant_id) is not None + ) + + +CHAINED_INVOKE_DIFFERENT_ACCOUNT_MESSAGE: Final[str] = ( + "Cannot start a CHAINED_INVOKE on a function in another account." +) +CHAINED_INVOKE_DIFFERENT_REGION_MESSAGE: Final[str] = ( + "Cannot start a CHAINED_INVOKE on a function in another region." +) + + +def chained_invoke_timeout_message(timeout_seconds: int) -> str: + """The service's message for a chained invoke that timed out.""" + return f"CHAINED_INVOKE timed out after {timeout_seconds} seconds." + + +def chained_invoke_timeout_error(timeout_seconds: int) -> ErrorObject: + """The error a timed-out CHAINED_INVOKE operation carries.""" + return ErrorObject( + message=chained_invoke_timeout_message(timeout_seconds), + type=CHAINED_INVOKE_TIMEOUT_ERROR_TYPE, + data=None, + stack_trace=None, + ) + + +def payload_size_bytes(payload: str) -> int: + """Size of ``payload`` as the service measures it: UTF-8 bytes.""" + return len(payload.encode("utf-8")) + + +@dataclass(frozen=True) +class FunctionTarget: + """A parsed chained-invoke target.""" + + name: str + qualifier: str | None + account_id: str | None + region: str | None + + +def parse_function_target(function_name: str) -> FunctionTarget: + """Parse a Lambda function name into its parts. + + Accepts the four forms Invoke accepts (bare name, name with + qualifier, partial ARN, full ARN). Raises ``ValueError`` for any + other string, including a full ARN missing its region or account. + """ + match = ( + _FUNCTION_NAME_PATTERN.fullmatch(function_name) + if 0 < len(function_name) <= MAX_FUNCTION_NAME_LENGTH + else None + ) + if match is None: + raise ValueError(function_name) + region: str | None = match.group("region") + account_id: str | None = match.group("account_id") + if match.group("arn_prefix") and (region is None or account_id is None): + raise ValueError(function_name) + return FunctionTarget( + name=match.group("name"), + qualifier=match.group("qualifier"), + account_id=account_id, + region=region, + ) + + +@dataclass(frozen=True) +class ChainedInvokeRequest: + """A chained-invoke dispatch request in wire form. + + ``function_name`` is the target's bare name and ``qualifier`` its + alias or version, if the parent gave one. + """ + + parent_execution_arn: str + operation_id: str + function_name: str + tenant_id: str | None + payload: str | None + account_id: str + trace_fields: dict | None = None + qualifier: str | None = None + + def lookup_keys(self) -> tuple[str, ...]: + """Keys the target's registration or configuration may sit under, + most specific first. See :func:`target_lookup_keys`.""" + return target_lookup_keys(self.function_name, self.qualifier) + + def child_qualifier(self) -> str: + """Qualifier the child execution records.""" + return self.qualifier if self.qualifier is not None else "$LATEST" + + def invoke_identifier(self) -> str: + """The function identifier to invoke, exactly as the parent wrote + it: ``name``, or ``name:qualifier`` including an explicit $LATEST.""" + if self.qualifier is None: + return self.function_name + return f"{self.function_name}:{self.qualifier}" + + +def target_lookup_keys(function_name: str, qualifier: str | None) -> tuple[str, ...]: + """Keys a target's registration or configuration may sit under, most + specific first. + + The runner has one registration per key, not versions: a qualified + key wins when present, otherwise the bare name serves every + qualifier. The key selects the registration only; what is invoked + is the identifier as the parent wrote it. + """ + if qualifier is None: + return (function_name,) + return (f"{function_name}:{qualifier}", function_name) + + +def invoke_identifier(function_name: str, qualifier: str) -> str: + """The identifier a durable execution's handler is invoked by: + ``name`` for $LATEST, else ``name:qualifier``.""" + if qualifier == "$LATEST": + return function_name + return f"{function_name}:{qualifier}" + + +@dataclass(frozen=True) +class ChildOutcome: + """Terminal outcome of a single-invocation chained target. + + ``timed_out`` marks an Invoke the endpoint did not answer within the + invocation timeout; ``error`` then carries the service's timeout + error and the operation ends TIMED_OUT rather than FAILED. + """ + + result: str | None = None + error: ErrorObject | None = None + timed_out: bool = False + + +def timed_out_outcome(timeout_seconds: int) -> ChildOutcome: + """Outcome of a plain target the endpoint did not answer for within ``timeout_seconds``.""" + return ChildOutcome( + error=chained_invoke_timeout_error(timeout_seconds), timed_out=True + ) + + +def function_timeout_outcome( + timeout_seconds: int, request_id: str, now: datetime +) -> ChildOutcome: + """Outcome of a plain target that ran past its function timeout. + + Lambda ends the function and answers the Invoke with a function + error, so the operation ends FAILED with that error, as it does at + the service. The message follows Lambda's own. + """ + stamp: str = now.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z" + return ChildOutcome( + error=ErrorObject( + message=f"{stamp} {request_id} Task timed out after {timeout_seconds:.2f} seconds", + type=FUNCTION_TIMEOUT_ERROR_TYPE, + data=None, + stack_trace=None, + ) + ) + + +def result_outcome(result: str | None) -> ChildOutcome: + """Outcome of a plain target that returned ``result``, size-checked.""" + if ( + result is not None + and payload_size_bytes(result) > MAX_CHAINED_INVOKE_PAYLOAD_BYTES + ): + return ChildOutcome( + error=ErrorObject.from_message(CHAINED_INVOKE_OUTPUT_TOO_LARGE_MESSAGE) + ) + return ChildOutcome(result=result) + + +@dataclass(frozen=True) +class StartChild: + """Start a new child durable execution. + + ``child_start.lambda_endpoint`` is ``None``. A per-execution endpoint + serves one function alone, and the child is another function, so + the parent's own endpoint is not the child's. The executor pins the + child to the endpoint the parent's chained targets go to + (``Invoker.inherit_endpoint``). + """ + + child_start: StartDurableExecutionInput + + +@dataclass(frozen=True) +class RunInvocation: + """Run the target as a single invocation off the worker lanes. + + The callable blocks for at most one invocation of a non-durable + target and returns that invocation's outcome. + """ + + invocation: Callable[[], ChildOutcome] + + +@dataclass(frozen=True) +class KnownOutcome: + """The terminal outcome is already known.""" + + outcome: ChildOutcome + + +DispatchResult = StartChild | RunInvocation | KnownOutcome + + +def failed_to_start(message: str, error_type: str | None = None) -> KnownOutcome: + """Build the outcome for a target that could not be dispatched. + + ``error_type`` is the Lambda API error code when one applies (for + example ``ResourceNotFoundException``); a dispatch failure with no + API counterpart carries a message only, as the service does. + """ + return KnownOutcome( + outcome=ChildOutcome( + error=ErrorObject( + message=message, + type=error_type, + data=None, + stack_trace=None, + ) + ) + ) + + +def function_not_found(function_name: str, hint: str = "") -> KnownOutcome: + """Build the outcome for a target that is not a known function.""" + return failed_to_start( + f"Function not found: {function_name}.{hint}", + error_type=FUNCTION_NOT_FOUND_ERROR_TYPE, + ) + + +def invoke_function( + lambda_client: Any, + function_name: str, + payload: str | None, + tenant_id: str | None, + invocation_timeout_seconds: int, +) -> ChildOutcome: + """Invoke non-durable ``function_name`` once at a Lambda-compatible endpoint. + + The response payload (or function error) is the outcome. An Invoke + API error fails the operation with that error's code and message, + as the service reports it. A read timeout means the target ran past + ``invocation_timeout_seconds``, which the service reports as a + chained-invoke timeout. + """ + try: + kwargs: dict[str, Any] = { + "FunctionName": function_name, + "InvocationType": "RequestResponse", + "Payload": payload if payload is not None else "{}", + } + if tenant_id is not None: + kwargs["TenantId"] = tenant_id + response: dict[str, Any] = lambda_client.invoke(**kwargs) + # The response headers arrive when the target returns; the body + # streams after them and can time out on its own. + raw_body: bytes = response["Payload"].read() + except ReadTimeoutError: + logger.info("Chained invoke of %s timed out", function_name) + return timed_out_outcome(invocation_timeout_seconds) + except ClientError as err: + logger.info("Chained invoke of %s failed to start: %s", function_name, err) + api_error: dict[str, Any] = err.response.get("Error", {}) + return ChildOutcome( + error=ErrorObject( + message=api_error.get("Message") or str(err), + type=api_error.get("Code"), + data=None, + stack_trace=None, + ) + ) + except Exception as err: # noqa: BLE001 — dispatch failure fails the operation + logger.info("Chained invoke of %s failed to start: %s", function_name, err) + return ChildOutcome( + error=ErrorObject( + message=f"Failed to invoke {function_name}: {err}", + type=None, + data=None, + stack_trace=None, + ) + ) + + try: + body: str = raw_body.decode("utf-8") + except UnicodeDecodeError: + return ChildOutcome( + error=ErrorObject.from_message(CHAINED_INVOKE_UTF8_DECODING_MESSAGE) + ) + if "FunctionError" in response: + return ChildOutcome(error=_error_from_invoke_body(body)) + return result_outcome(body if body else None) + + +class ChildDispatcher(Protocol): + """Dispatches a chained-invoke target for one environment.""" + + def preflight(self, target: FunctionTarget) -> ChildOutcome | None: + """The outcome of ``target`` known before anything is dispatched, + or ``None`` when dispatch may proceed. + + The service resolves a target before it schedules anything. A + target it cannot resolve comes back FAILED in the checkpoint + response itself, so the handler learns of it without + suspending. This is the runner's counterpart. It consults only + what the runner already holds (registrations or configurations) + and never the endpoint, so it returns at once. + """ + ... + + def dispatch(self, request: ChainedInvokeRequest) -> DispatchResult: + """Dispatch ``request`` and report how it will complete. + + Runs off the parent execution's worker lane. May block for the + duration of a single invocation of a non-durable target, but + must not block on a durable child reaching a terminal state. + """ + ... + + +@dataclass(frozen=True) +class RegisteredFunction: + """A function registered with the in-process runner.""" + + handler: Callable[..., Any] + is_durable: bool + execution_timeout_seconds: int = DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS + retention_period_days: int = DEFAULT_CHILD_RETENTION_PERIOD_DAYS + + +class FunctionRegistry: + """Name-to-handler registry for the in-process runner.""" + + def __init__(self) -> None: + self._functions: dict[str, RegisteredFunction] = {} + + def register( + self, + function_name: str, + handler: Callable[..., Any], + *, + is_durable: bool, + execution_timeout_seconds: int = DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS, + retention_period_days: int = DEFAULT_CHILD_RETENTION_PERIOD_DAYS, + ) -> None: + """Register ``handler`` under ``function_name``.""" + self._functions[function_name] = RegisteredFunction( + handler=handler, + is_durable=is_durable, + execution_timeout_seconds=execution_timeout_seconds, + retention_period_days=retention_period_days, + ) + + def get(self, function_name: str) -> RegisteredFunction | None: + """Return the registration for ``function_name`` if present.""" + return self._functions.get(function_name) + + +class InProcessChildDispatcher: + """Resolves chained-invoke targets from registered handlers. + + A durable registration becomes a child durable execution start; a + non-durable registration runs as a single handler call whose + payload and result cross the seam in serialized form; an unknown + name fails to start. + """ + + def __init__( + self, + registry: FunctionRegistry, + context_factory: Callable[[ChainedInvokeRequest], Any], + invocation_timeout_seconds: int, + clock: Clock, + ) -> None: + """``context_factory`` builds the Lambda context a plain target + receives; ``clock`` stamps a function-timeout error.""" + self._registry = registry + self._context_factory = context_factory + self._invocation_timeout_seconds = invocation_timeout_seconds + self._clock = clock + + def _find_registration( + self, function_name: str, qualifier: str | None + ) -> RegisteredFunction | None: + for key in target_lookup_keys(function_name, qualifier): + registration: RegisteredFunction | None = self._registry.get(key) + if registration is not None: + return registration + return None + + @staticmethod + def _not_found(function_name: str) -> KnownOutcome: + return function_not_found( + function_name, + hint=" Register it with register_durable_function or register_function.", + ) + + def preflight(self, target: FunctionTarget) -> ChildOutcome | None: + """Fail a target no handler is registered for.""" + if self._find_registration(target.name, target.qualifier) is None: + return self._not_found(target.name).outcome + return None + + def dispatch(self, request: ChainedInvokeRequest) -> DispatchResult: + """Dispatch ``request`` against the registered functions.""" + registration: RegisteredFunction | None = self._find_registration( + request.function_name, request.qualifier + ) + if registration is None: + return self._not_found(request.function_name) + + if registration.is_durable: + return StartChild( + child_start=StartDurableExecutionInput( + account_id=request.account_id, + function_name=request.function_name, + function_qualifier=request.child_qualifier(), + execution_name=str(uuid.uuid4()), + execution_timeout_seconds=registration.execution_timeout_seconds, + execution_retention_period_days=registration.retention_period_days, + invocation_id=None, + trace_fields=request.trace_fields, + tenant_id=request.tenant_id, + input=request.payload, + lambda_endpoint=None, + ) + ) + + handler: Callable[..., Any] = registration.handler + return RunInvocation( + invocation=lambda: self._run_single_invocation(handler, request) + ) + + def _run_single_invocation( + self, handler: Callable[..., Any], request: ChainedInvokeRequest + ) -> ChildOutcome: + """Run ``handler`` once, bounded by the invocation timeout. + + The handler runs on its own thread so the bound can be enforced. + A handler still running at the deadline is abandoned and the + outcome is Lambda's function-timeout error, as the runner does + for a handler invocation that exceeds the timeout. Python cannot + terminate a thread, so the abandoned handler keeps running until + it returns and its result is discarded. Lambda stops the sandbox + at the timeout; the in-process runner cannot, for this handler + as for its durable handlers. + """ + outcome: list[ChildOutcome] = [] + context: Any = self._context_factory(request) + + def run() -> None: + try: + payload: str | None = request.payload + event: Any = json.loads(payload) if payload else None + result: Any = handler(event, context) + outcome.append(result_outcome(json.dumps(result))) + except Exception as err: # noqa: BLE001 — the outcome carries the error + logger.info("Chained invoke target raised: %s", err) + outcome.append(ChildOutcome(error=ErrorObject.from_exception(err))) + + thread = threading.Thread( + target=run, name="durable-chained-invoke-target", daemon=True + ) + thread.start() + thread.join(self._invocation_timeout_seconds) + if thread.is_alive(): + logger.info("Chained invoke target exceeded the invocation timeout") + return function_timeout_outcome( + self._invocation_timeout_seconds, + str(getattr(context, "aws_request_id", "")), + self._clock.now(), + ) + return outcome[0] + + +def _error_from_invoke_body(body: str) -> ErrorObject: + """Build an ErrorObject from a function-error invoke response body.""" + try: + data: dict[str, Any] = json.loads(body) + except (json.JSONDecodeError, TypeError): + return ErrorObject.from_message(body or "Chained invoke failed.") + return ErrorObject( + message=data.get("errorMessage"), + type=data.get("errorType"), + data=data.get("errorData"), + stack_trace=data.get("stackTrace"), + ) + + +@dataclass(frozen=True) +class FunctionConfig: + """Durability configuration for a chained-invoke target function. + + Mirrors the Lambda function configuration: a function is durable + when it carries a ``DurableConfig``, whose ``ExecutionTimeout`` and + ``RetentionPeriodInDays`` bound its executions. A function without + one is a plain function. + """ + + is_durable: bool + execution_timeout_seconds: int = DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS + retention_period_days: int = DEFAULT_CHILD_RETENTION_PERIOD_DAYS + + @classmethod + def from_dict(cls, data: Any, name: str | None = None) -> FunctionConfig: + """Build from one entry of the function configurations. + + ``{}`` or ``null`` is a plain function. ``{"DurableConfig": {...}}`` + is a durable function; fields the ``DurableConfig`` omits take + the runner defaults. Any other shape raises ``ValueError`` naming + ``name``: a malformed entry must not turn a durable target plain + or fail with a bare traceback. + """ + if data is None: + return cls(is_durable=False) + if not isinstance(data, dict): + raise ValueError( + _malformed(name, "the entry must be a JSON object or null") + ) + durable_config: Any = data.get("DurableConfig") + if durable_config is None: + return cls(is_durable=False) + if not isinstance(durable_config, dict): + raise ValueError(_malformed(name, "DurableConfig must be a JSON object")) + return cls( + is_durable=True, + execution_timeout_seconds=_positive_int( + durable_config, + "ExecutionTimeout", + DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS, + name, + ), + retention_period_days=_positive_int( + durable_config, + "RetentionPeriodInDays", + DEFAULT_CHILD_RETENTION_PERIOD_DAYS, + name, + ), + ) + + +def _malformed(name: str | None, detail: str) -> str: + """Message for a malformed function configuration entry.""" + subject = ( + f"Function configuration for {name!r}" if name else "Function configuration" + ) + return f"{subject}: {detail}." + + +def _positive_int( + config: dict[str, Any], field: str, default: int, name: str | None +) -> int: + """``config[field]`` as a positive integer, or ``default`` when absent.""" + value: Any = config.get(field, default) + # bool is an int subclass; true/false are not durations. + if isinstance(value, bool) or not isinstance(value, int) or value < 1: + msg = _malformed(name, f"DurableConfig.{field} must be a positive integer") + raise ValueError(msg) + return value + + +FILE_URL_PREFIX: Final[str] = "file://" + + +@dataclass(frozen=True) +class FunctionConfigs: + """The functions a durable function may invoke, by name. + + Each key is a function name a parent may pass to ``context.invoke``; + each value is that function's :class:`FunctionConfig`. Built from + the ``--function-configs`` value with :meth:`from_value`, or from + an already parsed mapping with :meth:`from_dict`. + """ + + by_name: Mapping[str, FunctionConfig] + + @classmethod + def from_value(cls, value: str) -> FunctionConfigs: + """Parse the ``--function-configs`` value. + + ``value`` is a JSON object mapping function names to entries in + the shape :meth:`FunctionConfig.from_dict` reads, or + ``file://`` naming a file that holds one. Raises + ``ValueError`` for malformed JSON or a non-object, and ``OSError`` + for an unreadable file. + """ + text: str = value + if value.startswith(FILE_URL_PREFIX): + text = pathlib.Path(value[len(FILE_URL_PREFIX) :]).read_text( + encoding="utf-8" + ) + raw: Any = json.loads(text) + if not isinstance(raw, dict): + msg = "--function-configs must be a JSON object mapping function names to configurations." + raise ValueError(msg) + return cls.from_dict(raw) + + @classmethod + def from_dict(cls, raw: Mapping[str, Any]) -> FunctionConfigs: + """Build from a parsed name-to-entry mapping; a malformed entry raises ``ValueError``.""" + return cls( + {name: FunctionConfig.from_dict(entry, name) for name, entry in raw.items()} + ) + + def resolve(self, request: ChainedInvokeRequest) -> FunctionConfig | None: + """The configuration for ``request``'s target, if any.""" + return self.lookup(request.function_name, request.qualifier) + + def lookup( + self, function_name: str, qualifier: str | None + ) -> FunctionConfig | None: + """The configuration for a target, if any. + + The entry under the qualified identifier wins; otherwise the + bare name's entry serves every qualifier. + """ + for key in target_lookup_keys(function_name, qualifier): + config: FunctionConfig | None = self.by_name.get(key) + if config is not None: + return config + return None + + +class UnconfiguredChildDispatcher: + """Fails every chained invoke because no function configuration was given. + + The web runner cannot learn a target's durability from the endpoint, + so without the configuration it cannot dispatch; the error names the + option. + """ + + def preflight(self, target: FunctionTarget) -> ChildOutcome | None: + """Every target fails: the runner cannot resolve any.""" + return self._unconfigured(target.name).outcome + + def dispatch(self, request: ChainedInvokeRequest) -> DispatchResult: + """Fail ``request`` with a message naming the missing option.""" + return self._unconfigured(request.function_name) + + @staticmethod + def _unconfigured(function_name: str) -> KnownOutcome: + return failed_to_start( + f"Cannot invoke {function_name}: the local runner has no " + "function configurations. Start it with --function-configs " + "mapping each function name to its configuration." + ) + + +class EndpointChildDispatcher: + """Resolves chained-invoke targets against configured functions served + by a Lambda-compatible endpoint. + + A durable target becomes a child durable execution start: this + runner creates the child and invokes its handler at the endpoint by + function name through the normal invocation path. A non-durable + target runs as one RequestResponse invoke at the endpoint. An + unknown name fails to start with ``ResourceNotFoundException``. + """ + + def __init__( + self, + function_configs: FunctionConfigs, + client_provider: Callable[[str], Any], + invocation_timeout_seconds: int, + ) -> None: + """``client_provider`` maps a parent execution ARN to the Lambda + client its non-durable targets are invoked with, so a target goes + to the endpoint the parent is invoked at.""" + self._function_configs = function_configs + self._client_provider = client_provider + self._invocation_timeout_seconds = invocation_timeout_seconds + + def preflight(self, target: FunctionTarget) -> ChildOutcome | None: + """Fail a target no configuration names.""" + if self._function_configs.lookup(target.name, target.qualifier) is None: + return function_not_found(target.name).outcome + return None + + def dispatch(self, request: ChainedInvokeRequest) -> DispatchResult: + """Dispatch ``request`` against the configured functions.""" + config: FunctionConfig | None = self._function_configs.resolve(request) + if config is None: + return function_not_found(request.function_name) + + if config.is_durable: + return StartChild( + child_start=StartDurableExecutionInput( + account_id=request.account_id, + function_name=request.function_name, + function_qualifier=request.child_qualifier(), + execution_name=str(uuid.uuid4()), + execution_timeout_seconds=config.execution_timeout_seconds, + execution_retention_period_days=config.retention_period_days, + invocation_id=None, + trace_fields=request.trace_fields, + tenant_id=request.tenant_id, + input=request.payload, + lambda_endpoint=None, + ) + ) + + client: Any = self._client_provider(request.parent_execution_arn) + function_name: str = request.invoke_identifier() + payload: str | None = request.payload + tenant_id: str | None = request.tenant_id + timeout_seconds: int = self._invocation_timeout_seconds + return RunInvocation( + invocation=lambda: invoke_function( + client, function_name, payload, tenant_id, timeout_seconds + ) + ) diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/cli.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/cli.py index fce3d6ffa..b3b7c64d6 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/cli.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/cli.py @@ -25,6 +25,7 @@ from botocore.exceptions import ConnectionError # type: ignore +from aws_durable_execution_sdk_python_testing.child_dispatcher import FunctionConfigs from aws_durable_execution_sdk_python_testing.exceptions import ( DurableFunctionsLocalRunnerError, DurableFunctionsTestError, @@ -40,6 +41,14 @@ logger = logging.getLogger(__name__) +def _function_configs(value: str) -> FunctionConfigs: + """argparse type for ``--function-configs``: parse at the boundary.""" + try: + return FunctionConfigs.from_value(value) + except (ValueError, OSError) as exc: + raise argparse.ArgumentTypeError(str(exc)) from exc + + @dataclass(frozen=True) class CliConfig: """Configuration for the CLI application with environment variable support.""" @@ -219,7 +228,11 @@ def _create_start_server_parser(self, subparsers) -> None: "--invocation-timeout", type=int, default=900, - help="Per-invocation timeout in seconds, simulates Lambda Timeout (default: 900)", + help=( + "Per-invocation timeout in seconds, simulates Lambda Timeout " + "(default: 900). Also bounds how long the runner waits on any " + "Invoke it sends, with 60 seconds of headroom" + ), ) start_server_parser.add_argument( "--skip-time", @@ -227,6 +240,20 @@ def _create_start_server_parser(self, subparsers) -> None: default=False, help="Skip durable timer wall-clock waits; history keeps real modeled durations. Default is real timing (--no-skip-time); pass --skip-time to opt in", ) + start_server_parser.add_argument( + "--function-configs", + default=None, + type=_function_configs, + help=( + "The Lambda functions a durable function may invoke: a JSON " + "object mapping each function name to its configuration in " + "Lambda's own shape, or file:// to a file holding it. " + "A durable function carries a DurableConfig, e.g. " + '{"ProcessPayment": {"DurableConfig": {"ExecutionTimeout": 60, ' + '"RetentionPeriodInDays": 7}}, "LookupPrice": {}}. Required for ' + "chained invokes; without it every chained invoke fails" + ), + ) start_server_parser.set_defaults(func=self.start_server_command) def _create_invoke_parser(self, subparsers) -> None: @@ -301,6 +328,7 @@ def start_server_command(self, args: argparse.Namespace) -> int: store_path=args.store_path, invocation_timeout_seconds=args.invocation_timeout, skip_time=args.skip_time, + function_configs=args.function_configs, ) logger.info( diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py index e894caf1d..82880a1ba 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/execution.py @@ -1,5 +1,6 @@ from __future__ import annotations +import json import logging from dataclasses import dataclass, replace from datetime import datetime @@ -13,6 +14,7 @@ InvocationStatus, ) from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeDetails, ErrorObject, ExecutionDetails, Operation, @@ -29,6 +31,7 @@ # Import AWS exceptions from aws_durable_execution_sdk_python_testing.model import ( + executed_version, InvocationCompletedDetails, StartDurableExecutionInput, ) @@ -39,6 +42,16 @@ logger = logging.getLogger(__name__) +_CHAINED_INVOKE_TERMINAL_STATUSES: frozenset[OperationStatus] = frozenset( + { + OperationStatus.SUCCEEDED, + OperationStatus.FAILED, + OperationStatus.TIMED_OUT, + OperationStatus.STOPPED, + } +) + + class ExecutionStatus(Enum): """Execution status for API responses.""" @@ -121,6 +134,18 @@ def __init__( # Per-op payload size, tracked as a sidecar dict because # ``Operation`` is frozen upstream. self.operation_size_bytes: dict[str, int] = {} + # Chained-invoke linkage: operation_id -> the invoked child + # execution's ARN. Populated when a durable child execution is + # started for a CHAINED_INVOKE operation; absent for targets + # that run as a single invocation without their own execution. + self.chained_invoke_children: dict[str, str] = {} + # The parent execution's ARN when this execution was started by + # a chained invoke. A child's output is capped at the chained + # invoke limit, so validation needs to know. + self.parent_execution_arn: str | None = None + # The runner's region at creation; None when unknown (for + # example an execution built directly in a test). + self.region: str | None = None # Set when a trigger arrived while an invocation was already # in flight; the gate-release path consults it to decide whether # to schedule another invocation. @@ -211,6 +236,9 @@ def to_json_dict(self) -> dict[str, Any]: "HandlerSeenSeq": self.handler_seen_seq, "OperationLastTouchedSeq": dict(self.operation_last_touched_seq), "OperationSizeBytes": dict(self.operation_size_bytes), + "ChainedInvokeChildren": dict(self.chained_invoke_children), + "ParentExecutionArn": self.parent_execution_arn, + "Region": self.region, "NeedsReinvoke": self.needs_reinvoke, "LastCheckpoint": ( self.last_checkpoint.to_json_dict() if self.last_checkpoint else None @@ -264,6 +292,9 @@ def from_json_dict(cls, data: dict[str, Any]) -> Execution: data.get("OperationLastTouchedSeq", {}) ) execution.operation_size_bytes = dict(data.get("OperationSizeBytes", {})) + execution.chained_invoke_children = dict(data.get("ChainedInvokeChildren", {})) + execution.parent_execution_arn = data.get("ParentExecutionArn") + execution.region = data.get("Region") execution.needs_reinvoke = data.get("NeedsReinvoke", False) execution.current_invocation_id = data.get("CurrentInvocationId", "") last_checkpoint_data = data.get("LastCheckpoint") @@ -379,6 +410,17 @@ def has_pending_operations(self, execution: Execution) -> bool: return True return False + def has_unseen_changes(self) -> bool: + """True if any operation changed after the handler's last + observed sequence watermark.""" + return self.has_changes_after(self.handler_seen_seq) + + def has_changes_after(self, seq: int) -> bool: + """True if any operation changed after ``seq_counter`` was ``seq``.""" + return any( + touched > seq for touched in self.operation_last_touched_seq.values() + ) + def record_invocation_completion( self, start_timestamp: datetime, end_timestamp: datetime, request_id: str ) -> None: @@ -390,10 +432,23 @@ def record_invocation_completion( request_id=request_id, ) ) + + def mark_state_delivered(self) -> None: + """Reset the list of changed operations once an invocation's input + carries them. + + The service reports, on each invocation, the operations that + changed after the last state the handler observed. The runner + delivers state in the invocation input. So the list is reset when + that input is built, not when the invocation completes: an + operation that completes while the handler is still running is + reported on the next invocation. + """ self.updated_operation_ids = [] def _record_updated_operation(self, operation_id: str) -> None: - """Remember an operation changed outside the last invocation.""" + """Remember an operation that changed since the handler last saw + the state.""" if operation_id not in self.updated_operation_ids: self.updated_operation_ids.append(operation_id) @@ -646,6 +701,87 @@ def complete_callback_timeout( ) return self.operations[index] + def complete_chained_invoke( + self, + operation_id: str, + status: OperationStatus, + result: str | None = None, + error: ErrorObject | None = None, + now: datetime | None = None, + ) -> Operation: + """Transition a CHAINED_INVOKE operation to a terminal status. + + ``status`` must be SUCCEEDED, FAILED, TIMED_OUT, or STOPPED. + SUCCEEDED carries ``result``; the failure statuses carry + ``error``. + """ + index, operation = self.find_operation(operation_id) + self._require_chained_invoke(operation) + + if status not in _CHAINED_INVOKE_TERMINAL_STATUSES: + msg_bad_status: str = ( + f"Invalid terminal status for chained invoke: {status}" + ) + raise IllegalStateException(msg_bad_status) + if operation.status is not OperationStatus.STARTED: + msg_not_active: str = ( + f"Chained invoke operation [{operation_id}] is not active" + ) + raise IllegalStateException(msg_not_active) + if status is OperationStatus.SUCCEEDED and error is not None: + msg_success_error: str = ( + "Cannot provide an Error for a SUCCEEDED chained invoke." + ) + raise IllegalStateException(msg_success_error) + + with self._state_lock: + self.touch_operation(operation_id) + self.operations[index] = replace( + operation, + status=status, + end_timestamp=now if now is not None else real_now(), + chained_invoke_details=ChainedInvokeDetails(result=result, error=error), + ) + # State paging sizes a completed operation by its result or + # error, not by its START payload. + self.operation_size_bytes[operation_id] = len( + (result if result is not None else "").encode() + ) + (len(json.dumps(error.to_dict())) if error is not None else 0) + self._record_updated_operation(operation_id) + return self.operations[index] + + def executed_version(self) -> str: + """The version this execution runs, as the service reports it.""" + return executed_version(self.start_input.function_qualifier) + + def function_arn(self, default_region: str) -> str: + """The qualified function ARN of this execution. + + The service qualifies the ARN with the executed version, and + ``Version`` carries the same value. An execution stored before the + runner recorded regions has none; it reports ``default_region``. + """ + region: str = self.region if self.region is not None else default_region + return ( + f"arn:aws:lambda:{region}:{self.start_input.account_id}" + f":function:{self.start_input.function_name}:{self.executed_version()}" + ) + + def record_chained_invoke_child( + self, operation_id: str, child_execution_arn: str + ) -> None: + """Record the child execution ARN for a CHAINED_INVOKE operation.""" + with self._state_lock: + self.chained_invoke_children[operation_id] = child_execution_arn + + @staticmethod + def _require_chained_invoke(operation: Operation) -> None: + if operation.operation_type != OperationType.CHAINED_INVOKE: + msg: str = ( + f"Expected CHAINED_INVOKE operation, got {operation.operation_type}" + ) + raise IllegalStateException(msg) + def _end_execution( self, status: OperationStatus, now: datetime | None = None ) -> None: @@ -657,7 +793,7 @@ def _end_execution( # state change — record it via touch_operation so # introspection / GetDurableExecutionState reflects # it. The handler has returned by this point, so the - # touch is not load-bearing for a checkpoint delta. + # touch cannot change any checkpoint delta. self.touch_operation(execution_op.operation_id) self.operations[0] = replace( execution_op, diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/executor.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/executor.py index 3ace22b71..44a01bc01 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/executor.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/executor.py @@ -7,7 +7,7 @@ import threading import uuid from datetime import datetime -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, assert_never from aws_durable_execution_sdk_python.execution import ( DurableExecutionInvocationInput, @@ -30,6 +30,24 @@ from aws_durable_execution_sdk_python_testing.checkpoint.processor import ( DEFAULT_MAX_INVOCATION_PAGE_BYTES, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + CHILD_EXECUTION_OUTPUT_TOO_LARGE_MESSAGE, + FUNCTION_NOT_FOUND_ERROR_TYPE, + INVALID_FUNCTION_ARN_MESSAGE_FORMAT, + MAX_CHAINED_INVOKE_PAYLOAD_BYTES, + ChainedInvokeRequest, + ChildOutcome, + DispatchResult, + FunctionTarget, + KnownOutcome, + RunInvocation, + StartChild, + chained_invoke_timeout_error, + failed_to_start, + invoke_identifier, + parse_function_target, + payload_size_bytes, +) from aws_durable_execution_sdk_python_testing.clock import Clock, RealClock from aws_durable_execution_sdk_python_testing.checkpoint.transformer import ( CheckpointRequestDispatcher, @@ -41,6 +59,7 @@ ) from aws_durable_execution_sdk_python_testing.execution import ( Execution, + ExecutionStatus, OperationPaginatorState, ) from aws_durable_execution_sdk_python_testing.model import ( @@ -70,6 +89,7 @@ ExecutionObserver, apply_effects, ) +from aws_durable_execution_sdk_python_testing.threads import DaemonThreadPool from aws_durable_execution_sdk_python_testing.token import ( CallbackToken, CheckpointToken, @@ -88,6 +108,9 @@ from aws_durable_execution_sdk_python_testing.checkpoint.processor import ( CheckpointProcessor, ) + from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + ChildDispatcher, + ) from aws_durable_execution_sdk_python_testing.invoker import ( Invoker, InvokeResponse, @@ -99,6 +122,9 @@ class Executor(ExecutionObserver): + # Workers for the quick half of chained-invoke dispatch: resolving a + # target and creating a durable child. Each takes milliseconds. + QUICK_DISPATCH_WORKERS: int = 8 MAX_CONSECUTIVE_FAILED_ATTEMPTS: int = 5 RETRY_BACKOFF_SECONDS: int = 5 # GetDurableExecutionState page-count bounds, mirroring the service @@ -116,8 +142,14 @@ def __init__( invocation_timeout_seconds: int = 900, registry: ExecutionRegistry | None = None, clock: Clock | None = None, + child_dispatcher: ChildDispatcher | None = None, + region: str = "us-west-2", ): self._store = store + # The region this runner stands in for. Executions record it, so a + # chained-invoke target ARN in another region is rejected as the + # service rejects one. + self._region = region self._scheduler = scheduler self._invoker = invoker self._checkpoint_processor = checkpoint_processor @@ -126,7 +158,10 @@ def __init__( ) self._clock: Clock = clock if clock is not None else RealClock() self._invocation_timeout_seconds = invocation_timeout_seconds - self._dispatcher = CheckpointRequestDispatcher() + self._dispatcher = CheckpointRequestDispatcher( + chained_invoke_preflight=self.preflight_chained_invoke + ) + self._child_dispatcher = child_dispatcher self._max_invocation_page_bytes = ( max_invocation_page_bytes if max_invocation_page_bytes is not None @@ -144,15 +179,100 @@ def __init__( self._completion_events: dict[str, Event] = {} self._callback_timeouts: dict[str, Future] = {} self._callback_heartbeats: dict[str, Future] = {} - self._execution_timeout: Future | None = None + self._execution_timeouts: dict[str, Future] = {} + # Chained-invoke linkage: child execution ARN -> (parent + # execution ARN, operation id). Consulted on every terminal + # transition to complete the parent's CHAINED_INVOKE operation. + self._chained_invoke_links: dict[str, tuple[str, str]] = {} + self._links_lock: threading.Lock = threading.Lock() + # Dedicated pool for chained-invoke dispatch. Dispatching a + # non-durable target blocks for one invocation of it, so these + # calls stay off the scheduler's default executor, whose single + # thread every handler invocation depends on. Dispatching a + # durable target does not block: the child is launched and its + # terminal transition completes the parent's operation. + # Two pools, because the work has two durations. Resolving a + # target and creating a durable child take milliseconds. The + # Invoke of a plain target blocks until the target returns, up + # to the read timeout. Sharing one cap would let a wall of long + # Invokes delay a child start, which the service never does. + self._dispatch_pool: DaemonThreadPool = DaemonThreadPool( + max_workers=self.QUICK_DISPATCH_WORKERS, + thread_name_prefix="durable-chained-invoke", + ) + self._target_pool: DaemonThreadPool = DaemonThreadPool( + thread_name_prefix="durable-chained-invoke-target", + ) + # Dedicated pool for handler invocations, so concurrent + # executions (a parent and its chained children, or parallel + # web-driven executions) invoke their handlers in parallel. + self._invocation_pool: DaemonThreadPool = DaemonThreadPool( + thread_name_prefix="durable-invocation" + ) + # Set by shutdown(). A blocking call that returns after it + # must not touch the store or re-invoke anything. Child starts + # are counted under the condition, so shutdown can wait for the + # ones already running (milliseconds each) before it returns. + self._closing: bool = False + self._closing_cond = threading.Condition() + self._child_starts_in_flight: int = 0 + + def shutdown(self) -> None: + """Release executor-owned resources without waiting for blocked calls. + + Python cannot interrupt a thread blocked in a socket read. So a + handler invocation or plain-target Invoke in flight is left to + run until its endpoint answers or its read timeout expires. The + pools' workers are daemon threads, so neither this call nor + process exit waits for them. Queued work is cancelled and never + starts. A result that lands after this call is dropped by the + coroutine that awaited it. Inside a process that keeps running, + such a thread idles until its read timeout, holding one thread + and one socket. + """ + with self._closing_cond: + self._closing = True + # A child start already past its check finishes under the + # same condition; a new one sees the flag and does nothing. + while self._child_starts_in_flight > 0: + self._closing_cond.wait(timeout=5) + self._dispatch_pool.shutdown(wait=False, cancel_futures=True) + self._target_pool.shutdown(wait=False, cancel_futures=True) + self._invocation_pool.shutdown(wait=False, cancel_futures=True) + + @property + def region(self) -> str: + """The one region this runner emulates. + + A function has one region for its life, so the region is fixed + when the runner starts. Every execution is created in it and + reports it. + """ + return self._region def start_execution( self, input: StartDurableExecutionInput, # noqa: A002 ) -> StartDurableExecutionOutput: + execution = self._create_execution(input) + self._launch_execution(execution, input.execution_timeout_seconds) + return StartDurableExecutionOutput( + execution_arn=execution.durable_execution_arn + ) + + def _create_execution( + self, + input: StartDurableExecutionInput, # noqa: A002 + parent_execution_arn: str | None = None, + ) -> Execution: + """Create and persist a new execution without launching it. + + ``parent_execution_arn`` marks a child started by a chained + invoke; the child's output limit depends on it. + """ # Generate invocation_id if not provided if input.invocation_id is None: - input = StartDurableExecutionInput( + input = StartDurableExecutionInput( # noqa: A001 account_id=input.account_id, function_name=input.function_name, function_qualifier=input.function_qualifier, @@ -167,34 +287,36 @@ def start_execution( ) execution = Execution.new(input=input) + execution.region = self._region + execution.parent_execution_arn = parent_execution_arn execution.start(now=self._clock.now()) self._store.save(execution) logger.debug("Created execution with ARN: %s", execution.durable_execution_arn) + return execution + def _launch_execution(self, execution: Execution, timeout_seconds: int) -> None: + """Arm the execution's timeout and schedule its first invocation.""" + arn: str = execution.durable_execution_arn completion_event = self._scheduler.create_event() - self._completion_events[execution.durable_execution_arn] = completion_event + self._completion_events[arn] = completion_event # Schedule execution timeout - if input.execution_timeout_seconds > 0: + if timeout_seconds > 0: def timeout_handler(): error = ErrorObject.from_message( - f"Execution timed out after {input.execution_timeout_seconds} seconds." + f"Execution timed out after {timeout_seconds} seconds." ) - self.on_timed_out(execution.durable_execution_arn, error) + self.on_timed_out(arn, error) - self._execution_timeout = self._scheduler.call_later( + self._execution_timeouts[arn] = self._scheduler.call_later( timeout_handler, - delay=input.execution_timeout_seconds, + delay=timeout_seconds, completion_event=completion_event, ) # Schedule initial invocation to run immediately - self._invoke_execution(execution.durable_execution_arn) - - return StartDurableExecutionOutput( - execution_arn=execution.durable_execution_arn - ) + self._invoke_execution(arn) @staticmethod def _validate_execution_arn(execution_arn: str) -> None: @@ -262,7 +384,7 @@ def get_execution_details(self, execution_arn: str) -> GetDurableExecutionRespon return GetDurableExecutionResponse( durable_execution_arn=execution.durable_execution_arn, durable_execution_name=execution.start_input.execution_name, - function_arn=f"arn:aws:lambda:us-east-1:123456789012:function:{execution.start_input.function_name}", + function_arn=execution.function_arn(self._region), status=status, start_timestamp=execution_op.start_timestamp if execution_op.start_timestamp @@ -275,7 +397,7 @@ def get_execution_details(self, execution_arn: str) -> GetDurableExecutionRespon end_timestamp=execution_op.end_timestamp if execution_op.end_timestamp else None, - version="1.0", + version=execution.executed_version(), ) def list_executions( @@ -328,7 +450,9 @@ def list_executions( # Convert to ExecutionSummary objects execution_summaries: list[ExecutionSummary] = [ - ExecutionSummary.from_execution(execution, execution.current_status().value) + ExecutionSummary.from_execution( + execution, execution.current_status().value, self._region + ) for execution in executions ] @@ -656,6 +780,10 @@ def get_execution_history( op_update_ref: OperationUpdate | None = ( update if update.action in (OperationAction.RETRY, OperationAction.FAIL) + or ( + update.action == OperationAction.START + and update.operation_type == OperationType.CHAINED_INVOKE + ) else None ) @@ -667,12 +795,15 @@ def get_execution_history( execution.result, op_update_ref, include_execution_data, + child_execution_arn=execution.chained_invoke_children.get( + update.operation_id + ), ) if update.action == OperationAction.START: if update.operation_type == OperationType.CHAINED_INVOKE: all_events.append( - HistoryEvent.create_chained_invoke_event_pending(context) + HistoryEvent.create_chained_invoke_event_started(context) ) else: all_events.append(HistoryEvent.create_event_started(context)) @@ -724,9 +855,12 @@ def get_execution_history( execution.result, None, include_execution_data, + child_execution_arn=execution.chained_invoke_children.get( + op.operation_id + ), ) all_events.append( - HistoryEvent.create_chained_invoke_event_pending(context) + HistoryEvent.create_chained_invoke_event_started(context) ) if op.start_timestamp is not None: context = EventCreationContext( @@ -1260,9 +1394,13 @@ def _validate_invocation_response_and_store( execution_arn: str, response: DurableExecutionInvocationOutput, execution: Execution, + invocation_seq: int | None = None, ): """Validate response status and save it to the store if fine. + ``invocation_seq`` is the execution's ``seq_counter`` when the + invocation's input was built. + Raises: InvalidParameterValueException: If the response status is invalid. IllegalStateException: If the response status is valid but the execution is already completed. @@ -1293,13 +1431,40 @@ def _validate_invocation_response_and_store( "Cannot provide an Error for SUCCEEDED status." ) raise InvalidParameterValueException(msg_success_error) + if ( + execution.parent_execution_arn is not None + and response.result is not None + and payload_size_bytes(response.result) + > MAX_CHAINED_INVOKE_PAYLOAD_BYTES + ): + # A child returns its result to the parent's chained + # invoke, which caps it; the service applies the same + # bound to the invocation output as to a checkpoint. + logger.info("[%s] Child result exceeds the limit", execution_arn) + self._fail_workflow( + execution_arn, + ErrorObject.from_message( + CHILD_EXECUTION_OUTPUT_TOO_LARGE_MESSAGE + ), + ) + return logger.info("[%s] Execution succeeded", execution_arn) self._complete_workflow( execution_arn, result=response.result, error=None ) case InvocationStatus.PENDING: - if not execution.has_pending_operations(execution): + # An operation the handler waited on may complete between + # the handler's return and this check. A change the + # handler has not seen, after the invocation's input was + # built, earns a re-invoke, so PENDING is valid; only a + # handler that waited on nothing is in error. + if not execution.has_pending_operations(execution) and not ( + invocation_seq is not None + and execution.has_changes_after( + max(invocation_seq, execution.handler_seen_seq) + ) + ): msg_pending_ops: str = ( "Cannot return PENDING status with no pending operations." ) @@ -1314,13 +1479,18 @@ def _validate_invocation_response_and_store( def _begin_invocation( self, execution_arn: str - ) -> tuple[Execution, DurableExecutionInvocationInput] | None: + ) -> tuple[Execution, DurableExecutionInvocationInput, int] | None: """Claim the invocation gate and build the handler input. - Returns the execution and its invocation input when this call - claims the gate (PRE_INVOKE -> INVOKING); returns None when the - execution is already complete, the gate is COMPLETED, or another - invocation is in flight (recording a deferred re-invoke). + Returns the execution, its invocation input, and the execution's + ``seq_counter`` at that moment when this call claims the gate + (PRE_INVOKE -> INVOKING); returns None when the execution is + already complete, the gate is COMPLETED, or another invocation + is in flight (recording a deferred re-invoke). + + The counter is captured here, on the execution's serial lane, + because a completion queued behind this call mutates the same + object before the caller resumes. """ execution = self._store.load(execution_arn) @@ -1356,8 +1526,9 @@ def _begin_invocation( self._set_invocation_gate(execution_arn, InvocationState.INVOKING) execution.begin_new_invocation() invocation_input = self._invoker.create_invocation_input(execution=execution) + execution.mark_state_delivered() self._store.save(execution) - return execution, invocation_input + return execution, invocation_input, execution.seq_counter def _finish_invocation( self, @@ -1365,6 +1536,7 @@ def _finish_invocation( invoke_response: InvokeResponse, invocation_start: datetime, invocation_end: datetime, + invocation_seq: int | None = None, ) -> None: """Apply a completed handler invocation. @@ -1396,7 +1568,7 @@ def _finish_invocation( response = invoke_response.invocation_output try: self._validate_invocation_response_and_store( - execution_arn, response, execution + execution_arn, response, execution, invocation_seq ) except (InvalidParameterValueException, IllegalStateException) as e: logger.warning( @@ -1422,8 +1594,13 @@ def _finish_invocation( should_reinvoke = False else: self._set_invocation_gate(execution_arn, InvocationState.PRE_INVOKE) - should_reinvoke = reloaded.needs_reinvoke - if should_reinvoke: + # A deferred trigger is only actionable if something + # changed that the in-flight handler did not already + # observe through a checkpoint response. A completion the + # handler consumed mid-invocation earns no extra + # invocation. + should_reinvoke = reloaded.needs_reinvoke and reloaded.has_unseen_changes() + if reloaded.needs_reinvoke: reloaded.needs_reinvoke = False self._store.save(reloaded) @@ -1470,32 +1647,44 @@ async def invoke() -> None: # post-invoke processing run as separate steps: the two # state-mutation windows run on the execution's worker lane # (serialized with checkpoints), while the blocking invoke - # runs off the lane so the handler's checkpoints can run on - # it. + # runs on the invocation pool so concurrent executions + # (for example a chained child alongside its parent) invoke + # in parallel instead of queueing on one thread. + loop = asyncio.get_running_loop() try: - claim = await asyncio.to_thread( - lambda: self._registry.submit( + claim = await asyncio.wrap_future( + self._registry.submit( execution_arn, CallableTask(lambda: self._begin_invocation(execution_arn)), - ).result() + ) ) if claim is None: return - execution, invocation_input = claim + execution, invocation_input, invocation_seq = claim invocation_start = self._clock.now() invoke_response = await asyncio.wait_for( - asyncio.to_thread( - self._invoker.invoke, - execution.start_input.function_name, - invocation_input, - execution.start_input.lambda_endpoint, + loop.run_in_executor( + self._invocation_pool, + lambda: self._invoker.invoke( + invoke_identifier( + execution.start_input.function_name, + execution.start_input.function_qualifier, + ), + invocation_input, + execution.start_input.lambda_endpoint, + tenant_id=execution.start_input.tenant_id, + account_id=execution.start_input.account_id, + region_name=execution.region, + ), ), timeout=self._invocation_timeout_seconds, ) + if self._closing: + return invocation_end = self._clock.now() - await asyncio.to_thread( - lambda: self._registry.submit( + await asyncio.wrap_future( + self._registry.submit( execution_arn, CallableTask( lambda: self._finish_invocation( @@ -1503,26 +1692,38 @@ async def invoke() -> None: invoke_response, invocation_start, invocation_end, + invocation_seq, ) ), - ).result() + ) ) - except ResourceNotFoundException: + except ResourceNotFoundException as err: + if self._closing: + return logger.warning("[%s] Function No longer exists", execution_arn) - error_obj = ErrorObject.from_message(message="Function not found") - await asyncio.to_thread( - lambda: self._registry.submit( + # Keep the Lambda API error code, so a parent chained + # invoke sees ResourceNotFoundException on its operation. + error_obj = ErrorObject( + message=str(err) or "Function not found", + type=FUNCTION_NOT_FOUND_ERROR_TYPE, + data=None, + stack_trace=None, + ) + await asyncio.wrap_future( + self._registry.submit( execution_arn, CallableTask( lambda: self._fail_invocation_not_found( execution_arn, error_obj ) ), - ).result() + ) ) except asyncio.TimeoutError: + if self._closing: + return # Invocation killed by Lambda timeout. Step operations # stay in their current state (STARTED) — no checkpoint # was sent. Record the failed invocation and re-invoke. @@ -1535,8 +1736,8 @@ async def invoke() -> None: error_obj = ErrorObject.from_message( message=f"Function timed out after {self._invocation_timeout_seconds} seconds" ) - await asyncio.to_thread( - lambda: self._registry.submit( + await asyncio.wrap_future( + self._registry.submit( execution_arn, CallableTask( lambda: self._retry_after_timeout( @@ -1546,20 +1747,22 @@ async def invoke() -> None: invocation_end, ) ), - ).result() + ) ) except Exception as e: # noqa: BLE001 + if self._closing: + return # Handle invocation errors (network, function not found, etc.) logger.warning("[%s] Invocation failed: %s", execution_arn, e) error_obj = ErrorObject.from_exception(e) - await asyncio.to_thread( - lambda: self._registry.submit( + await asyncio.wrap_future( + self._registry.submit( execution_arn, CallableTask( lambda: self._retry_after_error(execution_arn, error_obj) ), - ).result() + ) ) return invoke @@ -1635,9 +1838,163 @@ def _complete_events(self, execution_arn: str): # complete doesn't actually checkpoint explicitly if event := self._completion_events.get(execution_arn): event.set() - if self._execution_timeout: - self._execution_timeout.cancel() - self._execution_timeout = None + if timeout_future := self._execution_timeouts.pop(execution_arn, None): + timeout_future.cancel() + self._notify_parent_of_terminal_child(execution_arn) + + def _notify_parent_of_terminal_child(self, child_arn: str) -> None: + """Complete the parent's CHAINED_INVOKE operation for a terminal child. + + The link from a child to its parent operation is looked up in + memory first. A runner restart empties that map. Both halves of + the link are persisted: the child's parent ARN and the parent's + map of operation id to child ARN. So a miss falls back to the + store, and a child that completes after a restart still reaches + its parent. No-op for an execution without a parent. The + parent-side transition runs on the parent's worker lane and the + parent is re-invoked to observe it. + """ + with self._links_lock: + link: tuple[str, str] | None = self._chained_invoke_links.pop( + child_arn, None + ) + child: Execution | None = None + try: + child = self._store.load(child_arn) + except Exception: # noqa: BLE001 — never let linkage failure kill the child's terminal path + logger.exception("[%s] Failed to read terminal execution", child_arn) + if link is None: + link = self._persisted_link(child) + if link is None: + return + parent_arn, operation_id = link + + if child is None: + status: OperationStatus = OperationStatus.FAILED + result: str | None = None + error: ErrorObject | None = ErrorObject.from_message( + "Chained invoke target completed but its outcome could not be read." + ) + else: + status, result, error = self._child_terminal_to_completion(child) + + self._registry.submit( + parent_arn, + CallableTask( + lambda: self._apply_chained_invoke_completion( + parent_arn, operation_id, status, result, error + ) + ), + ) + self._invoke_execution(parent_arn) + + def _persisted_link(self, child: Execution | None) -> tuple[str, str] | None: + """Rebuild a child's link to its parent operation from the store. + + Returns None when the child has no parent, when the parent cannot + be read, or when the parent no longer lists the child. + """ + if child is None or child.parent_execution_arn is None: + return None + try: + parent: Execution = self._store.load(child.parent_execution_arn) + except Exception: # noqa: BLE001 — a missing parent must not kill the child's terminal path + logger.exception( + "[%s] Failed to read parent %s", + child.durable_execution_arn, + child.parent_execution_arn, + ) + return None + for operation_id, child_arn in parent.chained_invoke_children.items(): + if child_arn == child.durable_execution_arn: + return (parent.durable_execution_arn, operation_id) + return None + + @staticmethod + def _child_terminal_to_completion( + child: Execution, + ) -> tuple[OperationStatus, str | None, ErrorObject | None]: + """Map a terminal child execution to the parent operation outcome.""" + result_output = child.result + status: ExecutionStatus = child.current_status() + match status: + case ExecutionStatus.SUCCEEDED: + return ( + OperationStatus.SUCCEEDED, + result_output.result if result_output else None, + None, + ) + case ExecutionStatus.TIMED_OUT: + # A timed-out child carries no payload; the service reports + # a fixed error type and the child's execution timeout. + return ( + OperationStatus.TIMED_OUT, + None, + chained_invoke_timeout_error( + child.start_input.execution_timeout_seconds + ), + ) + case ExecutionStatus.STOPPED: + # A stopped child surfaces the error object that the stop + # request carried, as the service does. + return ( + OperationStatus.STOPPED, + None, + result_output.error if result_output else None, + ) + case ExecutionStatus.FAILED: + error = ( + result_output.error + if result_output and result_output.error + else ErrorObject.from_message("Chained invoke failed.") + ) + return (OperationStatus.FAILED, None, error) + case ExecutionStatus.RUNNING: + msg: str = ( + f"Child execution {child.durable_execution_arn} has not " + "reached a terminal state." + ) + raise IllegalStateException(msg) + case _: + assert_never(status) + + def _apply_chained_invoke_completion( + self, + parent_arn: str, + operation_id: str, + status: OperationStatus, + result: str | None, + error: ErrorObject | None, + ) -> None: + """Stamp a chained-invoke outcome onto the parent operation. + + Runs on the parent's worker lane. Drops the outcome when the + parent is already terminal (its terminal state is authoritative) + or when the operation already completed. + """ + try: + execution = self.get_execution(parent_arn) + except ResourceNotFoundException: + logger.warning( + "[%s] Parent not found for chained invoke %s", parent_arn, operation_id + ) + return + if execution.is_complete: + return + _, operation = execution.find_operation(operation_id) + # Must agree with Execution.complete_chained_invoke, which + # accepts only STARTED; a raise here would land in a lane + # future nobody awaits. + if operation.status is not OperationStatus.STARTED: + return + execution.complete_chained_invoke( + operation_id, + status, + result=result, + error=error, + now=self._clock.now(), + ) + self._store.update(execution) def wait_until_complete( self, execution_arn: str, timeout: float | None = None @@ -1728,8 +2085,255 @@ def on_callback_created( # Schedule callback timeouts if configured self._schedule_callback_timeouts(execution_arn, callback_options, callback_id) + def on_chained_invoke_started( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> None: + """Dispatch a chained-invoke target. Observer method triggered by notifier. + + Scheduling only: the dispatch itself runs off the caller's + worker lane so the checkpoint that raised the effect returns + without waiting on the target. + """ + completion_event = self._completion_events.get(execution_arn) + self._scheduler.call_later( + self._dispatch_chained_invoke( + execution_arn, operation_id, function_name, tenant_id, payload + ), + completion_event=completion_event, + ) + # endregion ExecutionObserver + # region Chained invoke + def _dispatch_chained_invoke( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> Callable[[], Awaitable[None]]: + """Build the dispatch coroutine for a chained-invoke operation.""" + + async def dispatch() -> None: + loop = asyncio.get_running_loop() + try: + result: DispatchResult = await loop.run_in_executor( + self._dispatch_pool, + lambda: self._resolve_dispatch( + execution_arn, + operation_id, + function_name, + tenant_id, + payload, + ), + ) + if self._closing: + return + + match result: + case StartChild(child_start=child_start): + await loop.run_in_executor( + self._dispatch_pool, + lambda: self._start_child_execution( + execution_arn, operation_id, child_start + ), + ) + case RunInvocation(invocation=invocation): + invocation_result: ChildOutcome = await loop.run_in_executor( + self._target_pool, invocation + ) + if self._closing: + return + await self._finish_from_outcome( + execution_arn, operation_id, invocation_result + ) + case KnownOutcome(outcome=outcome): + await self._finish_from_outcome( + execution_arn, operation_id, outcome + ) + case _: + assert_never(result) + except Exception: + if self._closing: + # The pools refuse work after shutdown, and a target + # may raise after it. Neither outcome is applied. + return + logger.exception( + "[%s] Chained invoke dispatch failed for %s", + execution_arn, + operation_id, + ) + # An unexpected dispatch failure has no Lambda API error + # code, so the error carries a message only. + error = ErrorObject( + message="Chained invoke could not be dispatched.", + type=None, + data=None, + stack_trace=None, + ) + await self._finish_from_outcome( + execution_arn, operation_id, ChildOutcome(error=error) + ) + + return dispatch + + def preflight_chained_invoke(self, function_name: str) -> ErrorObject | None: + """Resolve a chained-invoke target at checkpoint time. + + Returns the error that fails the operation in the checkpoint + response, or ``None`` when the target can be dispatched. Runs on + the parent's worker lane inside the checkpoint, so it must not + block: the dispatcher answers from what it already holds. + """ + if self._child_dispatcher is None: + return failed_to_start( + "This runner has no chained-invoke dispatcher configured." + ).outcome.error + try: + target: FunctionTarget = parse_function_target(function_name) + except ValueError: + return failed_to_start( + INVALID_FUNCTION_ARN_MESSAGE_FORMAT.format(function_name) + ).outcome.error + outcome: ChildOutcome | None = self._child_dispatcher.preflight(target) + return outcome.error if outcome is not None else None + + def _resolve_dispatch( + self, + execution_arn: str, + operation_id: str, # noqa: ARG002 — part of the request identity in logs + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> DispatchResult: + """Resolve the dispatch result for a chained-invoke request.""" + if self._child_dispatcher is None: + return failed_to_start( + "This runner has no chained-invoke dispatcher configured." + ) + parent: Execution = self._store.load(execution_arn) + if parent.is_complete: + # The parent reached a terminal state between the checkpoint + # and this dispatch; its terminal state is authoritative. + return KnownOutcome(outcome=ChildOutcome()) + # Checkpoint validation already rejected a malformed target; this + # guards a dispatch that bypassed it. + try: + target: FunctionTarget = parse_function_target(function_name) + except ValueError: + return failed_to_start( + INVALID_FUNCTION_ARN_MESSAGE_FORMAT.format(function_name) + ) + request = ChainedInvokeRequest( + parent_execution_arn=execution_arn, + operation_id=operation_id, + function_name=target.name, + qualifier=target.qualifier, + tenant_id=tenant_id, + # An invoke without a payload delivers an empty JSON object + # to the target, matching the Lambda Invoke API. + payload=payload if payload is not None else "{}", + account_id=parent.start_input.account_id, + trace_fields=parent.start_input.trace_fields, + ) + return self._child_dispatcher.dispatch(request) + + async def _finish_from_outcome( + self, execution_arn: str, operation_id: str, outcome: ChildOutcome + ) -> None: + """Complete the operation from a known outcome and re-invoke.""" + status: OperationStatus + if outcome.timed_out: + status = OperationStatus.TIMED_OUT + elif outcome.error is None: + status = OperationStatus.SUCCEEDED + else: + status = OperationStatus.FAILED + await asyncio.wrap_future( + self._registry.submit( + execution_arn, + CallableTask( + lambda: self._apply_chained_invoke_completion( + execution_arn, + operation_id, + status, + outcome.result, + outcome.error, + ) + ), + ) + ) + self._invoke_execution(execution_arn) + + def _start_child_execution( + self, + parent_arn: str, + operation_id: str, + child_start: StartDurableExecutionInput, + ) -> None: + """Create, link, and launch a child durable execution. + + The linkage is registered and recorded on the parent before the + child's first invocation is scheduled, so a child that + completes immediately still finds its link. The child is pinned + to its parent's endpoint before that first invocation, so it is + invoked where the parent's chained targets go. + """ + # Shutdown can begin while this task runs. The check and the + # count are one step under the condition, and shutdown waits for + # the count to reach zero. So nothing is persisted or scheduled + # after shutdown: a child created then would never be invoked. + with self._closing_cond: + if self._closing: + return + self._child_starts_in_flight += 1 + try: + child: Execution = self._create_execution( + child_start, parent_execution_arn=parent_arn + ) + child_arn: str = child.durable_execution_arn + self._invoker.inherit_endpoint(child_arn, parent_arn) + with self._links_lock: + self._chained_invoke_links[child_arn] = (parent_arn, operation_id) + self._registry.submit( + parent_arn, + CallableTask( + lambda: self._record_child_on_parent( + parent_arn, operation_id, child_arn + ) + ), + ).result() + self._launch_execution(child, child_start.execution_timeout_seconds) + finally: + with self._closing_cond: + self._child_starts_in_flight -= 1 + self._closing_cond.notify_all() + + def _record_child_on_parent( + self, parent_arn: str, operation_id: str, child_arn: str + ) -> None: + """Record a child execution ARN on the parent's operation. + + Runs on the parent's worker lane. No-op when the parent is + already terminal. + """ + try: + execution = self.get_execution(parent_arn) + except ResourceNotFoundException: + return + if execution.is_complete: + return + execution.record_chained_invoke_child(operation_id, child_arn) + self._store.update(execution) + + # endregion Chained invoke + # region Callback Timeouts def _schedule_callback_timeouts( self, diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/invoker.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/invoker.py index d95ebc451..47e7ac5ed 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/invoker.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/invoker.py @@ -25,7 +25,10 @@ ResourceNotFoundException, ) from aws_durable_execution_sdk_python_testing.execution import OperationPaginatorState -from aws_durable_execution_sdk_python_testing.model import LambdaContext +from aws_durable_execution_sdk_python_testing.model import ( + LambdaContext, + executed_version, +) if TYPE_CHECKING: @@ -35,24 +38,66 @@ from aws_durable_execution_sdk_python_testing.execution import Execution -# Max Lambda function timeout is 15 minutes (900s); we give headroom for -# network round-trip and RIE startup. -_LAMBDA_READ_TIMEOUT_SECONDS = 960 -_LAMBDA_CLIENT_CONFIG = Config( - read_timeout=_LAMBDA_READ_TIMEOUT_SECONDS, - retries={"max_attempts": 0}, -) +# Every Invoke the runner sends is one Lambda invocation: a handler +# invocation, or a chained invoke of a non-durable target. The client +# waits for it to return, so the read timeout exceeds the emulated +# function timeout (``--invocation-timeout``, default 900 s) by a fixed +# headroom for the network round-trip and RIE startup. +DEFAULT_INVOCATION_TIMEOUT_SECONDS = 900 +LAMBDA_READ_TIMEOUT_HEADROOM_SECONDS = 60 + +def read_timeout_for(invocation_timeout_seconds: int) -> int: + """Client read timeout that outlasts one invocation of ``invocation_timeout_seconds``.""" + return invocation_timeout_seconds + LAMBDA_READ_TIMEOUT_HEADROOM_SECONDS -def create_lambda_client(endpoint_url: str | None, region_name: str) -> Any: - """Create a boto3 Lambda client configured for durable function invocations.""" - return boto3.client( +DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS = read_timeout_for( + DEFAULT_INVOCATION_TIMEOUT_SECONDS +) + +# Request header on every handler invocation. To a Lambda-compatible +# endpoint, a caller's Invoke of a durable function starts a new +# execution. A handler invocation is an Invoke of that function whose +# payload is a durable invocation input; the endpoint must run the +# handler once with it and return the handler's output. This header is +# how the runner says so. Endpoints that never start executions ignore it. +INVOCATION_MARKER_HEADER = "X-Dex-Handler-Invoke" +INVOCATION_MARKER_VALUE = "true" + + +def _add_invocation_marker(params: dict[str, Any], **_kwargs: Any) -> None: + """botocore ``before-call`` hook: mark the request as a handler invocation.""" + params.setdefault("headers", {})[INVOCATION_MARKER_HEADER] = INVOCATION_MARKER_VALUE + + +def create_lambda_client( + endpoint_url: str | None, + region_name: str, + read_timeout_seconds: int = DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + *, + mark_invocations: bool = True, +) -> Any: + """Create a boto3 Lambda client for the runner's Invoke calls. + + ``read_timeout_seconds`` bounds one Invoke. With ``mark_invocations`` + every Invoke carries :data:`INVOCATION_MARKER_HEADER`; the client + that invokes non-durable chained targets passes ``False`` so those + look like any caller's Invoke. + """ + + client: Any = boto3.client( "lambda", endpoint_url=endpoint_url, region_name=region_name, - config=_LAMBDA_CLIENT_CONFIG, + config=Config( + read_timeout=read_timeout_seconds, + retries={"max_attempts": 0}, + ), ) + if mark_invocations: + client.meta.events.register("before-call.lambda.Invoke", _add_invocation_marker) + return client @dataclass(frozen=True) @@ -63,7 +108,34 @@ class InvokeResponse: request_id: str -def create_test_lambda_context() -> LambdaContext: +# Values the in-process Lambda context reports when the caller gives none. +DEFAULT_TEST_REGION = "us-west-2" +DEFAULT_TEST_ACCOUNT_ID = "123456789012" +DEFAULT_TEST_FUNCTION_NAME = "test-function" + + +def create_test_lambda_context( + *, + region: str = DEFAULT_TEST_REGION, + account_id: str = DEFAULT_TEST_ACCOUNT_ID, + function_name: str = DEFAULT_TEST_FUNCTION_NAME, + tenant_id: str | None = None, +) -> LambdaContext: + """Build the Lambda context handed to an in-process handler. + + ``function_name`` is the identifier the function is invoked by: + ``name`` or ``name:qualifier``. Lambda fills ``function_name`` with + the bare name, ``function_version`` with the version that runs, and + ``invoked_function_arn`` with the ARN as invoked, so target code that + reads them sees the same values here. The runner keeps no versions: + a numeric qualifier is reported as the version, anything else + (no qualifier, ``$LATEST``, an alias) as ``$LATEST``. + + ``tenant_id`` is the execution's or invoke's tenant; ``None`` means + the invocation had none, as in Lambda. + """ + bare_name, _, qualifier = function_name.partition(":") + function_version: str = executed_version(qualifier) # Create client context as a dictionary, not as objects # LambdaContext.__init__ expects dictionaries and will create the objects internally client_context_dict = { @@ -87,8 +159,12 @@ def create_test_lambda_context() -> LambdaContext: aws_request_id="test-invoke-12345", client_context=client_context_dict, identity=cognito_identity_dict, - invoked_function_arn="arn:aws:lambda:us-west-2:123456789012:function:test-function", - tenant_id="test-tenant-789", + function_name=bare_name, + function_version=function_version, + invoked_function_arn=( + f"arn:aws:lambda:{region}:{account_id}:function:{function_name}" + ), + tenant_id=tenant_id, ) @@ -102,12 +178,19 @@ def invoke( function_name: str, input: DurableExecutionInvocationInput, endpoint_url: str | None = None, + tenant_id: str | None = None, + account_id: str | None = None, + region_name: str | None = None, ) -> InvokeResponse: ... # pragma: no cover def update_endpoint( self, endpoint_url: str, region_name: str ) -> None: ... # pragma: no cover + def inherit_endpoint( + self, child_execution_arn: str, parent_execution_arn: str + ) -> None: ... # pragma: no cover + class InProcessInvoker(Invoker): def __init__( @@ -115,10 +198,31 @@ def __init__( handler: Callable, service_client: InMemoryServiceClient, max_page_bytes: int = DEFAULT_MAX_INVOCATION_PAGE_BYTES, + region: str = DEFAULT_TEST_REGION, ): self.handler = handler + self._region = region self.service_client = service_client self._max_page_bytes = max_page_bytes + # Named handlers for chained-invoke targets. The root handler + # remains the fallback for the execution under test. + self._handlers: dict[str, Callable] = {} + + def register(self, function_name: str, handler: Callable) -> None: + """Register ``handler`` to be resolved by ``function_name``.""" + self._handlers[function_name] = handler + + def _resolve_handler(self, function_name: str) -> Callable: + """The handler for ``function_name``, which may carry a qualifier. + + A registration under the qualified identifier wins; otherwise + the bare name's registration serves every qualifier; otherwise + the runner's own handler. + """ + handler: Callable | None = self._handlers.get(function_name) + if handler is None: + handler = self._handlers.get(function_name.split(":", 1)[0]) + return handler if handler is not None else self.handler def create_invocation_input( self, execution: Execution @@ -139,16 +243,24 @@ def create_invocation_input( def invoke( self, - function_name: str, # noqa: ARG002 + function_name: str, input: DurableExecutionInvocationInput, endpoint_url: str | None = None, # noqa: ARG002 + tenant_id: str | None = None, + account_id: str | None = None, + region_name: str | None = None, # noqa: ARG002 — the context reports the runner's ) -> InvokeResponse: - # TODO: reasses if function_name will be used in future input_with_client = DurableExecutionInvocationInputWithClient.from_durable_execution_invocation_input( input, self.service_client ) - context = create_test_lambda_context() - response_dict = self.handler(input_with_client, context) + context = create_test_lambda_context( + region=self._region, + account_id=account_id or DEFAULT_TEST_ACCOUNT_ID, + function_name=function_name, + tenant_id=tenant_id, + ) + handler: Callable = self._resolve_handler(function_name) + response_dict = handler(input_with_client, context) output = DurableExecutionInvocationOutput.from_dict(response_dict) return InvokeResponse( invocation_output=output, request_id=context.aws_request_id @@ -157,20 +269,46 @@ def invoke( def update_endpoint(self, endpoint_url: str, region_name: str) -> None: """No-op for in-process invoker.""" + def inherit_endpoint( + self, child_execution_arn: str, parent_execution_arn: str + ) -> None: + """No-op for in-process invoker: there is no endpoint to inherit.""" + class LambdaInvoker(Invoker): def __init__( self, lambda_client: Any, max_page_bytes: int = DEFAULT_MAX_INVOCATION_PAGE_BYTES, + read_timeout_seconds: int = DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + endpoint_url: str = "", + region_name: str = "", ) -> None: + """``endpoint_url`` and ``region_name`` describe ``lambda_client``.""" self.lambda_client = lambda_client self._max_page_bytes = max_page_bytes - # Maps execution_arn -> endpoint for that execution - # Maps endpoint -> client to reuse clients across executions - self._execution_endpoints: dict[str, str] = {} - self._endpoint_clients: dict[str, Any] = {} - self._current_endpoint: str = "" # Track current endpoint for new executions + # Applied to every client this invoker creates for another endpoint. + self._read_timeout_seconds = read_timeout_seconds + # Clients are keyed by (endpoint URL, region): the region is the + # signing region, so the same URL in another region is another + # client. Marked clients invoke handlers; unmarked clients invoke + # non-durable chained targets. + self._endpoint_clients: dict[tuple[str, str], Any] = {} + self._unmarked_clients: dict[tuple[str, str], Any] = {} + # An execution without its own endpoint is pinned, at its first + # invocation, to the endpoint current at that moment. Every later + # handler invocation and every chained target it dispatches use + # the pinned endpoint, so an update_endpoint call while it runs + # does not split it across endpoints. A child it starts inherits + # the pin (inherit_endpoint). An execution with its own endpoint + # is pinned at its first chained dispatch instead: the pin then + # covers only its chained targets and children, and its own + # endpoint is signed in the execution's region. + self._execution_endpoints: dict[str, tuple[str, str]] = {} + # Endpoint and region for executions not yet pinned. + self._current: tuple[str, str] = (endpoint_url, region_name) + if endpoint_url: + self._endpoint_clients[self._current] = lambda_client self._lock = Lock() @staticmethod @@ -178,26 +316,79 @@ def create( endpoint_url: str, region_name: str, max_page_bytes: int = DEFAULT_MAX_INVOCATION_PAGE_BYTES, + read_timeout_seconds: int = DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, ) -> LambdaInvoker: """Create with the boto lambda client.""" - invoker = LambdaInvoker( - create_lambda_client(endpoint_url, region_name), + return LambdaInvoker( + create_lambda_client(endpoint_url, region_name, read_timeout_seconds), max_page_bytes=max_page_bytes, + read_timeout_seconds=read_timeout_seconds, + endpoint_url=endpoint_url, + region_name=region_name, ) - invoker._current_endpoint = endpoint_url - invoker._endpoint_clients[endpoint_url] = invoker.lambda_client - return invoker def update_endpoint(self, endpoint_url: str, region_name: str) -> None: - """Update the Lambda client endpoint.""" - # Cache client by endpoint to reuse across executions + """Update the Lambda endpoint and region for executions not yet pinned.""" + key: tuple[str, str] = (endpoint_url, region_name) with self._lock: - if endpoint_url not in self._endpoint_clients: - self._endpoint_clients[endpoint_url] = create_lambda_client( - endpoint_url, region_name + self.lambda_client = self._marked_client(key) + self._current = key + + def _marked_client(self, key: tuple[str, str]) -> Any: + """Client that marks its invokes as handler invocations. Caller holds the lock.""" + client: Any = self._endpoint_clients.get(key) + if client is None: + client = create_lambda_client( + key[0] or None, key[1], self._read_timeout_seconds + ) + self._endpoint_clients[key] = client + return client + + def _pinned_endpoint(self, durable_execution_arn: str) -> tuple[str, str]: + """The (endpoint, region) ``durable_execution_arn`` is pinned to, pinning it now if needed.""" + with self._lock: + pinned: tuple[str, str] | None = self._execution_endpoints.get( + durable_execution_arn + ) + if pinned is None: + pinned = self._current + self._execution_endpoints[durable_execution_arn] = pinned + return pinned + + def inherit_endpoint( + self, child_execution_arn: str, parent_execution_arn: str + ) -> None: + """Pin a chained child to the endpoint its parent's chained targets go to. + + The parent is pinned now if it has no pin yet. So the parent's + plain targets, the child's handler invocations, and the child's + own chained dispatches all use one endpoint and region, whatever + update_endpoint sets for executions started later. + """ + pinned: tuple[str, str] = self._pinned_endpoint(parent_execution_arn) + with self._lock: + self._execution_endpoints[child_execution_arn] = pinned + + def unmarked_client_for(self, durable_execution_arn: str) -> Any: + """Client for invoking the non-durable chained targets of one execution. + + The client targets the endpoint and region the execution is + pinned to, so a target goes where the execution's own handler + invocations go. It carries no invocation marker, so the endpoint + runs the target as any caller's Invoke. + """ + key: tuple[str, str] = self._pinned_endpoint(durable_execution_arn) + with self._lock: + client: Any = self._unmarked_clients.get(key) + if client is None: + client = create_lambda_client( + key[0] or None, + key[1], + self._read_timeout_seconds, + mark_invocations=False, ) - self.lambda_client = self._endpoint_clients[endpoint_url] - self._current_endpoint = endpoint_url + self._unmarked_clients[key] = client + return client def _get_client_for_execution( self, @@ -205,30 +396,28 @@ def _get_client_for_execution( lambda_endpoint: str | None = None, region_name: str | None = None, ) -> Any: - """Get the appropriate client for this execution.""" - # Use provided endpoint or fall back to cached endpoint for this execution + """Get the appropriate client for this execution. + + An execution with its own ``lambda_endpoint`` is invoked there, + signed in ``region_name``, the execution's region. That endpoint + serves the execution's function alone (under sam it is the + function's own container), so it is never recorded as the + execution's pin: the execution's chained targets are other + functions and go to the pinned endpoint, which routes by name. + Any other execution is invoked at the endpoint it is pinned to. + """ if lambda_endpoint: - if lambda_endpoint not in self._endpoint_clients: - self._endpoint_clients[lambda_endpoint] = create_lambda_client( - lambda_endpoint, region_name or "us-east-1" - ) - return self._endpoint_clients[lambda_endpoint] - - # Fallback to cached endpoint - if durable_execution_arn not in self._execution_endpoints: with self._lock: - if durable_execution_arn not in self._execution_endpoints: - self._execution_endpoints[durable_execution_arn] = ( - self._current_endpoint - ) - - endpoint = self._execution_endpoints[durable_execution_arn] + return self._marked_client( + (lambda_endpoint, region_name or "us-east-1") + ) - # If no endpoint configured, fall back to default client - if not endpoint: + key: tuple[str, str] = self._pinned_endpoint(durable_execution_arn) + if not key[0]: + # Built with a client and no endpoint: nothing else to pick. return self.lambda_client - - return self._endpoint_clients[endpoint] + with self._lock: + return self._marked_client(key) def create_invocation_input( self, execution: Execution @@ -250,13 +439,19 @@ def invoke( function_name: str, input: DurableExecutionInvocationInput, endpoint_url: str | None = None, + tenant_id: str | None = None, + account_id: str | None = None, # noqa: ARG002 — identity is the endpoint's + region_name: str | None = None, ) -> InvokeResponse: """Invoke AWS Lambda function and return durable execution result. Args: function_name: Name of the Lambda function to invoke input: Durable execution invocation input - endpoint_url: Lambda endpoint url + endpoint_url: The execution's own Lambda endpoint, if it has one + tenant_id: The execution's tenant, sent as the Invoke TenantId + account_id: The execution's account; unused, the endpoint owns identity + region_name: The execution's region; signs an Invoke at ``endpoint_url`` Returns: InvokeResponse: Response containing invocation output and request ID @@ -274,16 +469,20 @@ def invoke( # Get the client for this execution client = self._get_client_for_execution( - input.durable_execution_arn, endpoint_url + input.durable_execution_arn, endpoint_url, region_name ) + invoke_kwargs: dict[str, Any] = { + "FunctionName": function_name, + "InvocationType": "RequestResponse", # Synchronous invocation + "Payload": json.dumps(input.to_json_dict()), + } + if tenant_id is not None: + invoke_kwargs["TenantId"] = tenant_id + try: # Invoke AWS Lambda function using standard invoke method - response = client.invoke( - FunctionName=function_name, - InvocationType="RequestResponse", # Synchronous invocation - Payload=json.dumps(input.to_json_dict()), - ) + response = client.invoke(**invoke_kwargs) # Check HTTP status code status_code = response.get("StatusCode") diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py index 173fcf773..43026ec60 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/model.py @@ -279,6 +279,26 @@ def to_dict(self) -> dict[str, Any]: return result +LATEST_PUBLISHED_QUALIFIER = "$LATEST.PUBLISHED" + + +def executed_version(qualifier: str | None) -> str: + """The function version an invocation runs, as Lambda reports it. + + The service records the version Lambda executed: a numeric qualifier + names it, ``$LATEST.PUBLISHED`` is reported as itself, and an alias + resolves to the version it points at. The runner keeps no published + versions or aliases. So a numeric qualifier and ``$LATEST.PUBLISHED`` + are reported as given, and no qualifier, ``$LATEST`` or an alias is + reported as ``$LATEST``. + """ + if qualifier is not None and ( + qualifier.isdigit() or qualifier == LATEST_PUBLISHED_QUALIFIER + ): + return qualifier + return "$LATEST" + + @dataclass(frozen=True) class Execution: """Execution summary structure from Smithy model.""" @@ -317,14 +337,18 @@ def to_dict(self) -> dict[str, Any]: return result @classmethod - def from_execution(cls, execution, status: str) -> Execution: - """Create ExecutionSummary from Execution object.""" + def from_execution(cls, execution, status: str, default_region: str) -> Execution: + """Create ExecutionSummary from Execution object. + + ``default_region`` serves an execution stored before the runner + recorded regions; every other execution reports its own. + """ execution_op = execution.get_operation_execution_started() return cls( durable_execution_arn=execution.durable_execution_arn, durable_execution_name=execution.start_input.execution_name, - function_arn=f"arn:aws:lambda:us-east-1:123456789012:function:{execution.start_input.function_name}", + function_arn=execution.function_arn(default_region), status=status, start_timestamp=execution_op.start_timestamp if execution_op.start_timestamp @@ -1004,46 +1028,39 @@ def to_dict(self) -> dict[str, Any]: @dataclass(frozen=True) -class ChainedInvokePendingDetails: - """Chained Invoke Pending event details.""" +class ChainedInvokeStartedDetails: + """Chained invoke started event details. + + ExecutedVersion is omitted: the local runner performs no function + version resolution, so it has no resolved version to report. + """ - input: EventInput | None = None function_name: str | None = None + tenant_id: str | None = None + input: EventInput | None = None + durable_execution_arn: str | None = None @classmethod - def from_dict(cls, data: dict) -> ChainedInvokePendingDetails: + def from_dict(cls, data: dict) -> ChainedInvokeStartedDetails: input_data = None if input_dict := data.get("Input"): input_data = EventInput.from_dict(input_dict) return cls( - input=input_data, function_name=data.get("FunctionName"), + tenant_id=data.get("TenantId"), + input=input_data, + durable_execution_arn=data.get("DurableExecutionArn"), ) def to_dict(self) -> dict[str, Any]: result: dict[str, Any] = {} - if self.input is not None: - result["Input"] = self.input.to_dict() if self.function_name is not None: result["FunctionName"] = self.function_name - return result - - -@dataclass(frozen=True) -class ChainedInvokeStartedDetails: - """Chained invoke started event details.""" - - durable_execution_arn: str | None = None - - @classmethod - def from_dict(cls, data: dict) -> ChainedInvokeStartedDetails: - return cls( - durable_execution_arn=data.get("DurableExecutionArn"), - ) - - def to_dict(self) -> dict[str, Any]: - result: dict[str, Any] = {} + if self.tenant_id is not None: + result["TenantId"] = self.tenant_id + if self.input is not None: + result["Input"] = self.input.to_dict() if self.durable_execution_arn is not None: result["DurableExecutionArn"] = self.durable_execution_arn return result @@ -1288,6 +1305,7 @@ class EventCreationContext: durable_execution_invocation_output: DurableExecutionInvocationOutput | None = None operation_update: OperationUpdate | None = None include_execution_data: bool = False # noqa: FBT001, FBT002 + child_execution_arn: str | None = None @classmethod def create( @@ -1373,7 +1391,6 @@ class Event: step_started_details: StepStartedDetails | None = None step_succeeded_details: StepSucceededDetails | None = None step_failed_details: StepFailedDetails | None = None - chained_invoke_pending_details: ChainedInvokePendingDetails | None = None chained_invoke_started_details: ChainedInvokeStartedDetails | None = None chained_invoke_succeeded_details: ChainedInvokeSucceededDetails | None = None chained_invoke_failed_details: ChainedInvokeFailedDetails | None = None @@ -1448,12 +1465,6 @@ def from_dict(cls, data: dict) -> Event: if details_data := data.get("StepFailedDetails"): step_failed_details = StepFailedDetails.from_dict(details_data) - chained_invoke_pending_details = None - if details_data := data.get("ChainedInvokePendingDetails"): - chained_invoke_pending_details = ChainedInvokePendingDetails.from_dict( - details_data - ) - chained_invoke_started_details = None if details_data := data.get("ChainedInvokeStartedDetails"): chained_invoke_started_details = ChainedInvokeStartedDetails.from_dict( @@ -1530,7 +1541,6 @@ def from_dict(cls, data: dict) -> Event: step_started_details=step_started_details, step_succeeded_details=step_succeeded_details, step_failed_details=step_failed_details, - chained_invoke_pending_details=chained_invoke_pending_details, chained_invoke_started_details=chained_invoke_started_details, chained_invoke_succeeded_details=chained_invoke_succeeded_details, chained_invoke_failed_details=chained_invoke_failed_details, @@ -1589,10 +1599,6 @@ def to_dict(self) -> dict[str, Any]: result["StepSucceededDetails"] = self.step_succeeded_details.to_dict() if self.step_failed_details is not None: result["StepFailedDetails"] = self.step_failed_details.to_dict() - if self.chained_invoke_pending_details is not None: - result["ChainedInvokePendingDetails"] = ( - self.chained_invoke_pending_details.to_dict() - ) if self.chained_invoke_started_details is not None: result["ChainedInvokeStartedDetails"] = ( self.chained_invoke_started_details.to_dict() @@ -2010,31 +2016,20 @@ def create_step_event(cls, context: EventCreationContext) -> Event: # endregion step # region chained_invoke - @classmethod - def create_chained_invoke_event_pending( - cls, context: EventCreationContext - ) -> Event: - input: EventInput = EventInput.from_start_durable_execution_input( - context.start_durable_execution_input, context.include_execution_data - ) - return cls( - event_type=EventType.CHAINED_INVOKE_STARTED.value, - event_timestamp=context.start_timestamp, - sub_type=context.sub_type, - event_id=context.event_id, - operation_id=context.operation.operation_id, - name=context.operation.name, - parent_id=context.operation.parent_id, - chained_invoke_pending_details=ChainedInvokePendingDetails( - input=input, - function_name=context.start_durable_execution_input.function_name, - ), - ) - @classmethod def create_chained_invoke_event_started( cls, context: EventCreationContext ) -> Event: + update: OperationUpdate | None = context.operation_update + function_name: str | None = None + payload: str | None = None + if update is not None: + payload = update.payload + if update.chained_invoke_options is not None: + function_name = update.chained_invoke_options.function_name + # The API shape has TenantId here; the service's history does not + # return it for this event. The runner matches the service, so a + # history assertion that passes locally passes against it. return cls( event_type=EventType.CHAINED_INVOKE_STARTED.value, event_timestamp=context.start_timestamp, @@ -2044,7 +2039,16 @@ def create_chained_invoke_event_started( name=context.operation.name, parent_id=context.operation.parent_id, chained_invoke_started_details=ChainedInvokeStartedDetails( - durable_execution_arn=context.durable_execution_arn + function_name=function_name, + # No stored input, as for a target the runner could not + # resolve, means no Input on the event, as at the service. + input=EventInput( + payload=payload if context.include_execution_data else None, + truncated=not context.include_execution_data, + ) + if payload is not None + else None, + durable_execution_arn=context.child_execution_arn, ), ) @@ -2157,7 +2161,7 @@ def create_chained_invoke_event(cls, context: EventCreationContext) -> Event: """Create chained invoke event based on action.""" match context.operation.status: case OperationStatus.PENDING: - return cls.create_chained_invoke_event_pending(context) + return cls.create_chained_invoke_event_started(context) case OperationStatus.STARTED: return cls.create_chained_invoke_event_started(context) case OperationStatus.SUCCEEDED: @@ -2357,7 +2361,6 @@ def from_event_with_id(cls, event: Event, event_id: int) -> Event: step_started_details=event.step_started_details, step_succeeded_details=event.step_succeeded_details, step_failed_details=event.step_failed_details, - chained_invoke_pending_details=event.chained_invoke_pending_details, chained_invoke_started_details=event.chained_invoke_started_details, chained_invoke_succeeded_details=event.chained_invoke_succeeded_details, chained_invoke_failed_details=event.chained_invoke_failed_details, diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/observer.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/observer.py index 86ef8263a..52c216f8e 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/observer.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/observer.py @@ -15,6 +15,7 @@ from aws_durable_execution_sdk_python_testing.checkpoint.effects import ( CallbackCreated, + ChainedInvokeStarted, Completed, Failed, ) @@ -62,6 +63,17 @@ def on_callback_created( ) -> None: """Called when a callback is created.""" + @abstractmethod + def on_chained_invoke_started( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> None: + """Called when a chained invoke is accepted and must be dispatched.""" + class ExecutionNotifier: """Collects lifecycle effects raised while applying checkpoint updates. @@ -98,6 +110,25 @@ def notify_callback_created( ) ) + def notify_chained_invoke_started( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> None: + """Record that a chained invoke was accepted.""" + self.effects.append( + ChainedInvokeStarted( + execution_arn=execution_arn, + operation_id=operation_id, + function_name=function_name, + tenant_id=tenant_id, + payload=payload, + ) + ) + def apply_effects( effects: Iterable[CheckpointEffect], observer: ExecutionObserver @@ -115,3 +146,11 @@ def apply_effects( effect.callback_options, effect.callback_token, ) + elif isinstance(effect, ChainedInvokeStarted): + observer.on_chained_invoke_started( + effect.execution_arn, + effect.operation_id, + effect.function_name, + effect.tenant_id, + effect.payload, + ) diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/runner.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/runner.py index 28f719a08..a099652da 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/runner.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/runner.py @@ -44,12 +44,24 @@ InvalidParameterValueException, ResourceNotFoundException, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS, + DEFAULT_CHILD_RETENTION_PERIOD_DAYS, + ChildDispatcher, + EndpointChildDispatcher, + FunctionConfigs, + FunctionRegistry, + InProcessChildDispatcher, + UnconfiguredChildDispatcher, +) from aws_durable_execution_sdk_python_testing.executor import Executor from aws_durable_execution_sdk_python_testing.execution import ExecutionStatus from aws_durable_execution_sdk_python_testing.invoker import ( InProcessInvoker, LambdaInvoker, create_lambda_client, + create_test_lambda_context, + read_timeout_for, ) from aws_durable_execution_sdk_python_testing.model import ( GetDurableExecutionHistoryResponse, @@ -110,7 +122,9 @@ class WebRunnerConfig: store_type: StoreType = StoreType.MEMORY store_path: str | None = None # Path for filesystem store - # Timeout configuration + # Emulated Lambda function timeout, applied to each handler + # invocation. Also bounds how long the runner waits on any Invoke it + # sends (see invoker.read_timeout_for). invocation_timeout_seconds: int = 900 # Skip durable timer wall-clock waits (waits and step retries complete @@ -124,6 +138,11 @@ class WebRunnerConfig: # GetDurableExecutionState. None falls back to Executor default # (5 MB). max_invocation_page_bytes: int | None = None + # The functions a durable function may invoke, by name. Required for + # chained invokes: the runner cannot learn a target's durability from + # the Lambda endpoint. See FunctionConfigs.from_value for the input. + # Last, so positional construction from before it exists is unchanged. + function_configs: FunctionConfigs | None = None @dataclass(frozen=True) @@ -628,9 +647,13 @@ def __init__( invocation_timeout: int = 900, store: ExecutionStore | None = None, skip_time: bool = True, # noqa: FBT001, FBT002 + region: str = "us-west-2", ): self._execution_timeout = execution_timeout self._invocation_timeout = invocation_timeout + # Region the runner stands in for: reported in Lambda contexts and + # used to reject chained-invoke targets in another region. + self._region = region self._max_invocation_page_bytes = ( max_invocation_page_bytes if max_invocation_page_bytes is not None @@ -658,7 +681,9 @@ def __init__( handler, self._service_client, max_page_bytes=self._max_invocation_page_bytes, + region=region, ) + self._function_registry = FunctionRegistry() self._executor = Executor( store=self._store, scheduler=self._scheduler, @@ -668,10 +693,62 @@ def __init__( invocation_timeout_seconds=invocation_timeout, registry=self._registry, clock=self._clock, + child_dispatcher=InProcessChildDispatcher( + self._function_registry, + lambda request: create_test_lambda_context( + region=region, + account_id=request.account_id, + function_name=request.invoke_identifier(), + tenant_id=request.tenant_id, + ), + invocation_timeout_seconds=invocation_timeout, + clock=self._clock, + ), + region=region, ) # Wire up observer pattern - CheckpointProcessor uses this to notify executor of state changes self._checkpoint_processor.add_execution_observer(self._executor) + # A chained-invoke target the runner cannot resolve fails in the + # checkpoint response, as it does at the service. + self._checkpoint_processor.set_chained_invoke_preflight( + self._executor.preflight_chained_invoke + ) + + def register_durable_function( + self, + function_name: str, + handler: Callable, + execution_timeout: int = DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS, + retention_period_days: int = DEFAULT_CHILD_RETENTION_PERIOD_DAYS, + ) -> DurableFunctionTestRunner: + """Register a durable function as a chained-invoke target. + + A ``ctx.invoke`` of ``function_name`` runs ``handler`` as its + own durable execution with the given timeout and retention. + Returns the runner for chaining. + """ + self._invoker.register(function_name, handler) + self._function_registry.register( + function_name, + handler, + is_durable=True, + execution_timeout_seconds=execution_timeout, + retention_period_days=retention_period_days, + ) + return self + + def register_function( + self, function_name: str, handler: Callable + ) -> DurableFunctionTestRunner: + """Register a non-durable function as a chained-invoke target. + + A ``ctx.invoke`` of ``function_name`` calls ``handler`` once + with the deserialized payload and a test Lambda context. + Returns the runner for chaining. + """ + self._function_registry.register(function_name, handler, is_durable=False) + return self def __enter__(self): return self @@ -680,6 +757,7 @@ def __exit__(self, exc_type, exc_val, exc_tb): self.close() def close(self): + self._executor.shutdown() self._registry.shutdown() self._scheduler.stop() @@ -691,6 +769,7 @@ def run( function_name: str = "test-function", execution_name: str = "execution-name", account_id: str = "123456789012", + tenant_id: str | None = None, ) -> DurableFunctionTestResult: if timeout is not None and execution_timeout is None: warnings.warn( @@ -710,6 +789,7 @@ def run( function_name=function_name, execution_name=execution_name, account_id=account_id, + tenant_id=tenant_id, ) return self.wait_for_result( @@ -729,6 +809,21 @@ def send_callback_failure( def send_callback_heartbeat(self, callback_id: str) -> None: self._executor.send_callback_heartbeat(callback_id=callback_id) + def get_execution_history( + self, + execution_arn: str, + include_execution_data: bool = False, # noqa: FBT001, FBT002 + ) -> GetDurableExecutionHistoryResponse: + """Return the event history for an execution. + + Works for the execution under test and for any chained child + execution (the child's ARN is on the parent's + ChainedInvokeStarted event). + """ + return self._executor.get_execution_history( + execution_arn, include_execution_data=include_execution_data + ) + def run_async( self, input: str | None = None, # noqa: A002 @@ -737,6 +832,7 @@ def run_async( function_name: str = "test-function", execution_name: str = "execution-name", account_id: str = "123456789012", + tenant_id: str | None = None, ) -> str: if timeout is not None and execution_timeout is None: execution_timeout = timeout @@ -754,7 +850,7 @@ def run_async( execution_retention_period_days=7, invocation_id="inv-12345678-1234-1234-1234-123456789012", trace_fields={"trace_id": "abc123", "span_id": "def456"}, - tenant_id="tenant-001", + tenant_id=tenant_id, input=input, ) @@ -829,7 +925,22 @@ def handler(event: Any, context: DurableContext): # noqa: ARG001 class WebRunner: - """Web server runner for durable functions testing with HTTP API endpoints.""" + """Web server runner for durable functions testing with HTTP API endpoints. + + Handlers run at the Lambda endpoint (``WebRunnerConfig.lambda_endpoint``); + the runner drives each execution by invoking its handler there. + + Chained invokes (``context.invoke``) need ``WebRunnerConfig.function_configs``, + the targets by name (:class:`FunctionConfigs`). Without it every + chained invoke fails with an error naming the option. A durable + target runs as a child execution the runner drives; a plain target is + one RequestResponse Invoke at the endpoint, bounded by the invocation + timeout, whose function error fails the operation. Every handler + invocation the runner sends carries ``X-Dex-Handler-Invoke: true`` + (:data:`invoker.INVOCATION_MARKER_HEADER`), so an endpoint that starts + an execution on Invoke of a durable function runs the handler once + instead; plain targets are sent without it. + """ def __init__(self, config: WebRunnerConfig) -> None: """Initialize WebRunner with configuration. @@ -893,9 +1004,16 @@ def start(self) -> None: if self._config.max_invocation_page_bytes is not None else DEFAULT_MAX_INVOCATION_PAGE_BYTES ) + read_timeout_seconds: int = read_timeout_for( + self._config.invocation_timeout_seconds + ) + lambda_client: Any = self._create_boto3_client(read_timeout_seconds) self._invoker = LambdaInvoker( - self._create_boto3_client(), + lambda_client, max_page_bytes=resolved_max_page_bytes, + read_timeout_seconds=read_timeout_seconds, + endpoint_url=self._config.lambda_endpoint, + region_name=self._config.local_runner_region, ) # Create shared CheckpointProcessor @@ -908,6 +1026,19 @@ def start(self) -> None: clock=clock, ) + # Non-durable targets are invoked with the invoker's unmarked client + # for the parent's endpoint, so the endpoint runs them as any + # caller's Invoke and they go where the parent's invocations go. + child_dispatcher: ChildDispatcher + if self._config.function_configs is not None: + child_dispatcher = EndpointChildDispatcher( + self._config.function_configs, + self._invoker.unmarked_client_for, + invocation_timeout_seconds=self._config.invocation_timeout_seconds, + ) + else: + child_dispatcher = UnconfiguredChildDispatcher() + # Create executor with all dependencies including checkpoint processor self._executor = Executor( store=self._store, @@ -918,10 +1049,15 @@ def start(self) -> None: invocation_timeout_seconds=self._config.invocation_timeout_seconds, registry=self._registry, clock=clock, + child_dispatcher=child_dispatcher, + region=self._config.local_runner_region, ) # Add executor as observer to the checkpoint processor checkpoint_processor.add_execution_observer(self._executor) + checkpoint_processor.set_chained_invoke_preflight( + self._executor.preflight_chained_invoke + ) # Start the scheduler self._scheduler.start() @@ -964,6 +1100,12 @@ def stop(self) -> None: self._server = None + if self._executor is not None: + try: + self._executor.shutdown() + except Exception: + logger.exception("error shutting down executor") + if self._registry is not None: try: self._registry.shutdown() @@ -982,10 +1124,12 @@ def stop(self) -> None: self._invoker = None self._executor = None - def _create_boto3_client(self) -> Any: + def _create_boto3_client(self, read_timeout_seconds: int) -> Any: """Create boto3 client for Lambda service. - Creates a boto3 client with the local runner endpoint and region from configuration. + Creates a boto3 client with the local runner endpoint and region from + configuration. ``read_timeout_seconds`` is passed through to + :func:`create_lambda_client`. Returns: Configured boto3 client for Lambda service @@ -997,6 +1141,7 @@ def _create_boto3_client(self) -> Any: return create_lambda_client( endpoint_url=self._config.lambda_endpoint, region_name=self._config.local_runner_region, + read_timeout_seconds=read_timeout_seconds, ) diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/threads.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/threads.py new file mode 100644 index 000000000..fa199cd39 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/threads.py @@ -0,0 +1,157 @@ +"""A bounded pool of daemon threads for blocking calls. + +The runner runs two kinds of blocking call off the scheduler loop: a +handler invocation, and the synchronous Invoke of a plain chained-invoke +target. Each blocks until the endpoint answers or the read timeout +expires. Python cannot interrupt a thread blocked in a socket read. So +the runner cannot stop such a call. It can stop the call from holding +the process. + +``concurrent.futures.ThreadPoolExecutor`` does not allow that. Its +workers are not daemon threads, and the interpreter joins them at exit. +So a runner closed while one Invoke is blocked keeps the process alive +until that read timeout. This pool has the same shape (a worker cap, a +FIFO queue, lazily started workers) and differs in two ways. Its workers +are daemon threads, which the interpreter does not wait for. Its +``shutdown`` drops queued work, refuses new work, and never joins. +""" + +from __future__ import annotations + +import os +import queue +import threading +from concurrent.futures import Executor, Future +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from collections.abc import Callable + + +def default_max_workers() -> int: + """The worker cap ``ThreadPoolExecutor`` uses when none is given.""" + return min(32, (os.cpu_count() or 1) + 4) + + +class _WorkItem: + def __init__( + self, + future: Future[Any], + fn: Callable[..., Any], + args: tuple[Any, ...], + kwargs: dict[str, Any], + ) -> None: + self.future = future + self.fn = fn + self.args = args + self.kwargs = kwargs + + def run(self) -> None: + if not self.future.set_running_or_notify_cancel(): + return + try: + result = self.fn(*self.args, **self.kwargs) + except BaseException as exc: # noqa: BLE001 — the future carries the error + self.future.set_exception(exc) + else: + self.future.set_result(result) + + +class _Stop: + """Sentinel that makes an idle worker return.""" + + +class DaemonThreadPool(Executor): + """Run callables on at most ``max_workers`` daemon threads, in FIFO order. + + A worker is started when work is submitted and no worker is idle, + until the cap is reached. Beyond the cap, work waits in the queue in + submission order, as in ``ThreadPoolExecutor``. + + ``shutdown`` marks the pool closed, cancels every queued future, and + tells idle workers to return. A worker busy in a blocking call is + left alone: it finishes when its call returns, and its result lands + on a future nobody reads. ``wait`` is accepted for API compatibility + and ignored, because joining a blocked worker would wait for the read + timeout this pool exists to avoid. Workers are daemon threads, so + process exit does not wait for them either. + """ + + def __init__( + self, max_workers: int | None = None, thread_name_prefix: str = "" + ) -> None: + if max_workers is None: + max_workers = default_max_workers() + if max_workers <= 0: + msg = "max_workers must be greater than 0" + raise ValueError(msg) + self._max_workers = max_workers + self._thread_name_prefix = thread_name_prefix or "daemon-pool" + self._queue: queue.SimpleQueue[_WorkItem | _Stop] = queue.SimpleQueue() + self._idle = threading.Semaphore(0) + self._threads: list[threading.Thread] = [] + self._closed = False + self._lock = threading.Lock() + + @property + def max_workers(self) -> int: + return self._max_workers + + def submit( + self, fn: Callable[..., Any], /, *args: Any, **kwargs: Any + ) -> Future[Any]: + with self._lock: + if self._closed: + msg = "cannot schedule new futures after shutdown" + raise RuntimeError(msg) + future: Future[Any] = Future() + self._queue.put(_WorkItem(future, fn, args, kwargs)) + self._adjust_thread_count() + return future + + def _adjust_thread_count(self) -> None: + # An idle worker will take the item; do not start another. + if self._idle.acquire(blocking=False): + return + if len(self._threads) >= self._max_workers: + return + thread = threading.Thread( + target=self._worker, + name=f"{self._thread_name_prefix}_{len(self._threads)}", + daemon=True, + ) + self._threads.append(thread) + thread.start() + + def _worker(self) -> None: + while True: + item: _WorkItem | _Stop = self._queue.get() + if isinstance(item, _Stop): + return + item.run() + del item + self._idle.release() + + def shutdown(self, wait: bool = True, *, cancel_futures: bool = False) -> None: # noqa: ARG002, FBT001, FBT002 + """Close the pool without waiting for busy workers. + + Queued work is cancelled whatever ``cancel_futures`` says: work + that has not started must not start after shutdown. ``wait`` is + ignored, see the class docstring. + """ + with self._lock: + if self._closed: + return + self._closed = True + drained: list[_WorkItem] = [] + while True: + try: + item = self._queue.get_nowait() + except queue.Empty: + break + if isinstance(item, _WorkItem): + drained.append(item) + for _ in self._threads: + self._queue.put(_Stop()) + for item in drained: + item.future.cancel() diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/web/handlers.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/web/handlers.py index e5048eea6..8cddfaa1f 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/web/handlers.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/web/handlers.py @@ -809,17 +809,28 @@ def handle(self, parsed_route: Route, request: HTTPRequest) -> HTTPResponse: # try: body = self._parse_json_body(request) endpoint_url = body.get("EndpointUrl") - region_name = body.get("RegionName", "us-east-1") + region_name = body.get("RegionName") if not endpoint_url: return self._handle_aws_exception( InvalidParameterValueException("EndpointUrl is required") ) + # The runner emulates one region, fixed at startup, and every + # execution is created in it. A request that names another + # region would give the runner two regions at once, so it is + # rejected instead of applied to the invoker alone. + if region_name is not None and region_name != self.executor.region: + return self._handle_aws_exception( + InvalidParameterValueException( + f"RegionName must be {self.executor.region}, " + f"the region this runner emulates; got {region_name}" + ) + ) # Update the invoker's Lambda endpoint invoker = self.executor._invoker # noqa: SLF001 logger.info("Updating lambda endpoint to %s", endpoint_url) - invoker.update_endpoint(endpoint_url, region_name) + invoker.update_endpoint(endpoint_url, self.executor.region) return self._success_response( {"message": "Lambda endpoint updated successfully"} ) diff --git a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/worker/registry.py b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/worker/registry.py index 7170dd42b..cd8edcbbb 100644 --- a/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/worker/registry.py +++ b/packages/aws-durable-execution-sdk-python-testing/src/aws_durable_execution_sdk_python_testing/worker/registry.py @@ -34,6 +34,7 @@ def __init__(self, store: ExecutionStore, scheduler: Scheduler) -> None: self._scheduler = scheduler self._workers: dict[str, ExecutionWorker] = {} self._lock = threading.Lock() + self._shut_down: bool = False # Submitting can lose a race with a worker tearing its own lane # down (it stops the lane, then leaves the registry), so a @@ -44,8 +45,16 @@ def __init__(self, store: ExecutionStore, scheduler: Scheduler) -> None: _MAX_SUBMIT_ATTEMPTS: int = 5 def get_or_create(self, execution_arn: str) -> ExecutionWorker: - """Return the worker for ``execution_arn``, creating it if absent.""" + """Return the worker for ``execution_arn``, creating it if absent. + + Raises: + RuntimeError: After ``shutdown``. A lane created then would + outlive the runner that owned it. + """ with self._lock: + if self._shut_down: + msg: str = "execution registry is shut down" + raise RuntimeError(msg) worker: ExecutionWorker | None = self._workers.get(execution_arn) if worker is None: worker = ExecutionWorker.create( @@ -101,6 +110,7 @@ def shutdown(self) -> None: terminal status still has a live lane, which this stops. """ with self._lock: + self._shut_down = True workers: list[ExecutionWorker] = list(self._workers.values()) self._workers.clear() for worker in workers: diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/chained_invoke_endpoint_int_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/chained_invoke_endpoint_int_test.py new file mode 100644 index 000000000..258da4086 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/chained_invoke_endpoint_int_test.py @@ -0,0 +1,218 @@ +"""Chained invoke through the endpoint dispatcher, with a scripted Lambda endpoint. + +The web runner resolves a chained-invoke target from the function +configuration file and invokes handlers at a Lambda endpoint. These tests +replace that endpoint with a scripted :class:`Invoker`, so the executor, +checkpoint processor, and dispatcher run for real while the endpoint's +behavior is fixed by the test. +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING, Any +from unittest.mock import Mock + +from aws_durable_execution_sdk_python.execution import ( + DurableExecutionInvocationInput, + DurableExecutionInvocationOutput, + InvocationStatus, +) +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, + OperationAction, + OperationStatus, + OperationType, + OperationUpdate, +) + +from aws_durable_execution_sdk_python_testing.checkpoint.processor import ( + CheckpointProcessor, +) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + EndpointChildDispatcher, + FunctionConfig, + FunctionConfigs, +) +from aws_durable_execution_sdk_python_testing.exceptions import ( + ResourceNotFoundException, +) +from aws_durable_execution_sdk_python_testing.executor import Executor +from aws_durable_execution_sdk_python_testing.invoker import ( + InvokeResponse, + LambdaInvoker, +) +from aws_durable_execution_sdk_python_testing.model import StartDurableExecutionInput +from aws_durable_execution_sdk_python_testing.scheduler import Scheduler +from aws_durable_execution_sdk_python_testing.stores.memory import ( + InMemoryExecutionStore, +) + +if TYPE_CHECKING: + from aws_durable_execution_sdk_python_testing.execution import Execution + +PARENT = "parent" +CHILD = "child" + + +class _ScriptedEndpoint: + """An :class:`Invoker` that plays a Lambda endpoint. + + ``parent`` is a durable handler that invokes ``child`` once and then + returns the chained-invoke operation's outcome. ``child`` does not + exist at the endpoint, so invoking it raises ResourceNotFoundException + exactly as :class:`LambdaInvoker` does for a 404 from Lambda. + """ + + def __init__(self) -> None: + self.executor: Executor | None = None + self._paginator = LambdaInvoker(Mock()) + # child ARN -> parent ARN, as pinned by the executor. + self.inherited: dict[str, str] = {} + + def create_invocation_input( + self, execution: Execution + ) -> DurableExecutionInvocationInput: + return self._paginator.create_invocation_input(execution) + + def update_endpoint(self, endpoint_url: str, region_name: str) -> None: + msg = f"the scripted endpoint cannot move to {endpoint_url} ({region_name})" + raise AssertionError(msg) + + def inherit_endpoint( + self, child_execution_arn: str, parent_execution_arn: str + ) -> None: + self.inherited[child_execution_arn] = parent_execution_arn + + def invoke( + self, + function_name: str, + input: DurableExecutionInvocationInput, # noqa: A002 + endpoint_url: str | None = None, # noqa: ARG002 + tenant_id: str | None = None, # noqa: ARG002 + account_id: str | None = None, # noqa: ARG002 + region_name: str | None = None, # noqa: ARG002 + ) -> InvokeResponse: + if function_name == CHILD: + # The executor pins a child to its parent's endpoint before + # scheduling the child's first invocation. + assert input.durable_execution_arn in self.inherited + msg = f"Function not found: {CHILD}" + raise ResourceNotFoundException(msg) + assert function_name == PARENT + assert self.executor is not None + operations = input.initial_execution_state.operations + invoke_op = next( + ( + op + for op in operations + if op.operation_type == OperationType.CHAINED_INVOKE + ), + None, + ) + if invoke_op is None: + self.executor.checkpoint_execution( + input.durable_execution_arn, + input.checkpoint_token, + [ + OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=OperationAction.START, + name="call-child", + payload="{}", + chained_invoke_options=ChainedInvokeOptions( + function_name=CHILD, tenant_id=None + ), + ) + ], + ) + return _output(InvocationStatus.PENDING) + if invoke_op.status == OperationStatus.STARTED: + return _output(InvocationStatus.PENDING) + details = invoke_op.chained_invoke_details + assert details is not None + error = details.error + return _output( + InvocationStatus.SUCCEEDED, + result=json.dumps( + { + "status": invoke_op.status.value, + "type": error.type if error else None, + "message": error.message if error else None, + } + ), + ) + + +def _output(status: InvocationStatus, result: str | None = None) -> InvokeResponse: + return InvokeResponse( + invocation_output=DurableExecutionInvocationOutput( + status=status, result=result + ), + request_id="scripted", + ) + + +def test_configured_durable_target_missing_at_endpoint_keeps_lambda_error_code(): + """The file says ``child`` is durable, but the endpoint has no such function. + + The child execution's first invocation fails with + ResourceNotFoundException, and the parent's operation carries that + error code, as it does when the service cannot find the function. + """ + endpoint = _ScriptedEndpoint() + store = InMemoryExecutionStore() + scheduler = Scheduler() + checkpoint_processor = CheckpointProcessor(store, scheduler) + executor = Executor( + store=store, + scheduler=scheduler, + invoker=endpoint, + checkpoint_processor=checkpoint_processor, + child_dispatcher=EndpointChildDispatcher( + FunctionConfigs( + {CHILD: FunctionConfig(is_durable=True, execution_timeout_seconds=30)} + ), + client_provider=lambda _arn: Mock(), + invocation_timeout_seconds=5, + ), + ) + checkpoint_processor.add_execution_observer(executor) + endpoint.executor = executor + scheduler.start() + try: + output = executor.start_execution( + StartDurableExecutionInput( + account_id="123456789012", + function_name=PARENT, + function_qualifier="$LATEST", + execution_name="run", + execution_timeout_seconds=30, + execution_retention_period_days=1, + input="{}", + ) + ) + parent_arn: str = output.execution_arn or "" + assert executor.wait_until_complete(parent_arn, timeout=20) + + parent = executor.get_execution(parent_arn) + assert parent.result is not None + outcome: dict[str, Any] = json.loads(parent.result.result or "{}") + assert outcome == { + "status": "FAILED", + "type": "ResourceNotFoundException", + "message": "Function not found: child", + } + # The child execution exists and failed with the same error. + child_arn = parent.chained_invoke_children["invoke-1"] + child = executor.get_execution(child_arn) + assert child.parent_execution_arn == parent_arn + assert child.result is not None + assert child.result.error is not None + assert child.result.error.type == "ResourceNotFoundException" + # The executor pinned the child to its parent's endpoint. + assert endpoint.inherited == {child_arn: parent_arn} + finally: + executor.shutdown() + scheduler.stop() diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/processors/invoke_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/processors/invoke_test.py new file mode 100644 index 000000000..3c8b010a4 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/processors/invoke_test.py @@ -0,0 +1,170 @@ +"""Tests for the CHAINED_INVOKE operation processor.""" + +from datetime import UTC, datetime + +import pytest +from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, + ErrorObject, + OperationAction, + OperationStatus, + OperationSubType, + OperationType, + OperationUpdate, +) + +from aws_durable_execution_sdk_python_testing.checkpoint.effects import ( + ChainedInvokeStarted, +) +from aws_durable_execution_sdk_python_testing.checkpoint.processors.invoke import ( + ChainedInvokeProcessor, +) +from aws_durable_execution_sdk_python_testing.exceptions import ( + InvalidParameterValueException, +) +from aws_durable_execution_sdk_python_testing.observer import ExecutionNotifier + +NOW = datetime(2024, 1, 1, tzinfo=UTC) + + +def _start_update( + options: ChainedInvokeOptions | None, + payload: str | None = '{"n": 1}', +) -> OperationUpdate: + return OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + sub_type=OperationSubType.CHAINED_INVOKE, + action=OperationAction.START, + name="double", + payload=payload, + chained_invoke_options=options, + ) + + +def test_start_records_started_operation_and_started_effect(): + notifier = ExecutionNotifier() + update = _start_update( + ChainedInvokeOptions(function_name="child", tenant_id="tenant-a") + ) + + operation = ChainedInvokeProcessor().process( + update, None, notifier, "arn-parent", NOW + ) + + # STARTED at checkpoint, as the service's checkpoint response carries + # a resolvable target; no other status is visible before completion. + assert operation.status == OperationStatus.STARTED + assert operation.operation_type == OperationType.CHAINED_INVOKE + assert operation.start_timestamp == NOW + assert operation.end_timestamp is None + assert operation.chained_invoke_details is not None + assert operation.chained_invoke_details.result is None + assert operation.chained_invoke_details.error is None + + assert notifier.effects == [ + ChainedInvokeStarted( + execution_arn="arn-parent", + operation_id="invoke-1", + function_name="child", + tenant_id="tenant-a", + payload='{"n": 1}', + ) + ] + + +def test_start_with_a_passing_preflight_dispatches(): + seen: list[str] = [] + + def preflight(function_name: str) -> ErrorObject | None: + seen.append(function_name) + return None + + notifier = ExecutionNotifier() + operation = ChainedInvokeProcessor(preflight).process( + _start_update(ChainedInvokeOptions(function_name="child:prod", tenant_id=None)), + None, + notifier, + "arn-parent", + NOW, + ) + + assert seen == ["child:prod"] # the identifier as the handler wrote it + assert operation.status == OperationStatus.STARTED + assert len(notifier.effects) == 1 + + +def test_start_with_a_failing_preflight_fails_in_the_checkpoint_response(): + """The service returns a target it cannot resolve as FAILED in the + checkpoint response, so the handler raises without suspending. The + processor records the failure at once and dispatches nothing.""" + error = ErrorObject( + message="Function not found: child.", + type="ResourceNotFoundException", + data=None, + stack_trace=None, + ) + notifier = ExecutionNotifier() + + operation = ChainedInvokeProcessor(lambda _name: error).process( + _start_update(ChainedInvokeOptions(function_name="child", tenant_id=None)), + None, + notifier, + "arn-parent", + NOW, + ) + + assert operation.status == OperationStatus.FAILED + assert operation.start_timestamp == NOW + assert operation.end_timestamp == NOW + assert operation.chained_invoke_details is not None + assert operation.chained_invoke_details.error == error + assert operation.chained_invoke_details.result is None + assert notifier.effects == [] + + +def test_start_without_options_is_rejected(): + with pytest.raises(InvalidParameterValueException) as excinfo: + ChainedInvokeProcessor().process( + _start_update(options=None), None, ExecutionNotifier(), "arn", NOW + ) + assert "requires ChainedInvokeOptions" in str(excinfo.value) + + +@pytest.mark.parametrize( + "action", + [OperationAction.CANCEL, OperationAction.SUCCEED, OperationAction.FAIL], +) +def test_non_start_action_is_rejected(action: OperationAction): + update = OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=action, + chained_invoke_options=ChainedInvokeOptions(function_name="child"), + ) + with pytest.raises(InvalidParameterValueException) as excinfo: + ChainedInvokeProcessor().process(update, None, ExecutionNotifier(), "arn", NOW) + assert str(excinfo.value) == "Invalid action for CHAINED_INVOKE operation." + + +def test_a_failed_preflight_stores_the_update_without_its_input(): + """The service keeps no input for a target it cannot resolve. The + stored update keeps the function name and drops the payload, so + history shows the name and the error and the input weighs nothing.""" + error = ErrorObject.from_message("Function not found: child.") + processor = ChainedInvokeProcessor(lambda _name: error) + update = _start_update(ChainedInvokeOptions(function_name="child", tenant_id=None)) + failed = processor.process(update, None, ExecutionNotifier(), "arn-parent", NOW) + + stored = processor.stored_update(update, None, failed) + + assert stored.payload is None + assert stored.chained_invoke_options.function_name == "child" + + +def test_a_started_invoke_stores_the_update_as_sent(): + processor = ChainedInvokeProcessor(lambda _name: None) + update = _start_update(ChainedInvokeOptions(function_name="child", tenant_id=None)) + started = processor.process(update, None, ExecutionNotifier(), "arn-parent", NOW) + + assert processor.stored_update(update, None, started) is update diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/transformer_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/transformer_test.py index f152b0c52..fff7218ac 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/transformer_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/transformer_test.py @@ -13,6 +13,7 @@ import pytest from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, ErrorObject, OperationAction, OperationType, @@ -23,6 +24,9 @@ from aws_durable_execution_sdk_python_testing.checkpoint.processors.base import ( OperationProcessor, ) +from aws_durable_execution_sdk_python_testing.checkpoint.processors.invoke import ( + ChainedInvokeProcessor, +) from aws_durable_execution_sdk_python_testing.checkpoint.transformer import ( CheckpointRequestDispatcher, ) @@ -77,6 +81,21 @@ def test_dispatcher_init_accepts_custom_processors(): assert dispatcher.processors is custom +def test_dispatcher_init_with_a_preflight_replaces_only_the_invoke_processor(): + def preflight(_name: str): + return None + + dispatcher = CheckpointRequestDispatcher(chained_invoke_preflight=preflight) + + invoke_processor = dispatcher.processors[OperationType.CHAINED_INVOKE] + assert isinstance(invoke_processor, ChainedInvokeProcessor) + assert invoke_processor._preflight is preflight # noqa: SLF001 + # The shared default table is untouched. + defaults = CheckpointRequestDispatcher().processors + assert defaults[OperationType.CHAINED_INVOKE] is not invoke_processor + assert dispatcher.processors[OperationType.STEP] is defaults[OperationType.STEP] + + def test_apply_updates_with_empty_list_is_a_noop(): dispatcher = CheckpointRequestDispatcher() execution = _make_execution() @@ -470,3 +489,52 @@ def test_apply_updates_records_size_for_bytes_payload(): # bytes payload length == 12. assert execution.operation_size_bytes["with-bytes"] == 12 + + +def _invoke_start(payload: str) -> OperationUpdate: + return OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=OperationAction.START, + payload=payload, + chained_invoke_options=ChainedInvokeOptions( + function_name="child", tenant_id=None + ), + ) + + +def test_a_failed_preflight_stores_no_input_and_weighs_nothing(): + """The service keeps no input for a target it cannot resolve. The + stored update has no payload and the operation's size is zero; the + function name stays for history.""" + error = ErrorObject.from_message("Function not found: child.") + dispatcher = CheckpointRequestDispatcher(chained_invoke_preflight=lambda _n: error) + execution = _make_execution() + + dispatcher.apply_updates( + execution, + [_invoke_start('{"secret": "retain-me"}')], + None, + lambda _id: None, + datetime.now(UTC), + ) + + assert execution.updates[-1].payload is None + assert execution.updates[-1].chained_invoke_options.function_name == "child" + assert execution.operation_size_bytes["invoke-1"] == 0 + + +def test_a_started_invoke_stores_its_input_and_size(): + dispatcher = CheckpointRequestDispatcher(chained_invoke_preflight=lambda _n: None) + execution = _make_execution() + + dispatcher.apply_updates( + execution, + [_invoke_start('{"secret": "keep-me"}')], + None, + lambda _id: None, + datetime.now(UTC), + ) + + assert execution.updates[-1].payload == '{"secret": "keep-me"}' + assert execution.operation_size_bytes["invoke-1"] == len('{"secret": "keep-me"}') diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/checkpoint_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/checkpoint_test.py index 548e98833..956ee34ed 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/checkpoint_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/checkpoint_test.py @@ -4,6 +4,7 @@ import pytest from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, ErrorObject, Operation, OperationAction, @@ -16,6 +17,9 @@ MAX_ERROR_PAYLOAD_SIZE_BYTES, CheckpointValidator, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + MAX_CHAINED_INVOKE_PAYLOAD_BYTES, +) from aws_durable_execution_sdk_python_testing.exceptions import ( InvalidParameterValueException, ) @@ -166,6 +170,71 @@ def test_validate_payload_sizes_error_within_limit(): CheckpointValidator.validate_input(updates, execution) +def _chained_invoke_start(payload: str) -> OperationUpdate: + return OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=OperationAction.START, + payload=payload, + chained_invoke_options=ChainedInvokeOptions( + function_name="child", tenant_id=None + ), + ) + + +def test_chained_invoke_input_at_the_limit_is_accepted(): + execution = _create_test_execution() + CheckpointValidator.validate_input( + [_chained_invoke_start("x" * MAX_CHAINED_INVOKE_PAYLOAD_BYTES)], execution + ) + + +def test_chained_invoke_input_over_the_limit_is_rejected(): + """The input limit counts UTF-8 bytes, so a multi-byte payload trips it sooner.""" + execution = _create_test_execution() + over_by_chars = "x" * (MAX_CHAINED_INVOKE_PAYLOAD_BYTES + 1) + over_by_bytes = "\u00e9" * (MAX_CHAINED_INVOKE_PAYLOAD_BYTES // 2 + 1) + for payload in (over_by_chars, over_by_bytes): + with pytest.raises(InvalidParameterValueException) as exc_info: + CheckpointValidator.validate_input( + [_chained_invoke_start(payload)], execution + ) + assert str(exc_info.value) == ( + "CHAINED_INVOKE input payload size must be less than or equal to " + "1048576 bytes." + ) + + +def _execution_succeed(payload: str) -> OperationUpdate: + return OperationUpdate( + operation_id="exec-1", + operation_type=OperationType.EXECUTION, + action=OperationAction.SUCCEED, + payload=payload, + ) + + +def test_child_execution_output_over_the_limit_is_rejected(): + """A child of a chained invoke returns at most 1 MiB to its parent.""" + execution = _create_test_execution() + execution.parent_execution_arn = "parent-arn" + with pytest.raises(InvalidParameterValueException) as exc_info: + CheckpointValidator.validate_input( + [_execution_succeed("x" * (MAX_CHAINED_INVOKE_PAYLOAD_BYTES + 1))], + execution, + ) + assert str(exc_info.value) == ( + "Execution output payload size must be less than or equal to 1048576 bytes." + ) + + +def test_top_level_execution_output_is_not_bound_by_the_chained_invoke_limit(): + execution = _create_test_execution() + CheckpointValidator.validate_input( + [_execution_succeed("x" * (MAX_CHAINED_INVOKE_PAYLOAD_BYTES + 1))], execution + ) + + def test_validate_duplicate_operation_ids(): """Test validation allows duplicate operation IDs in same batch. @@ -371,7 +440,8 @@ def test_validate_operation_status_transition_invoke(): action=OperationAction.CANCEL, ) ] - CheckpointValidator.validate_input(updates, execution) + with pytest.raises(InvalidParameterValueException): + CheckpointValidator.validate_input(updates, execution) def test_validate_operation_status_transition_execution(): diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/operations/invoke_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/operations/invoke_test.py index 3300b1ab0..3dd3c6cf2 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/operations/invoke_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/operations/invoke_test.py @@ -2,6 +2,7 @@ import pytest from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, Operation, OperationAction, OperationStatus, @@ -15,16 +16,134 @@ from aws_durable_execution_sdk_python_testing.exceptions import ( InvalidParameterValueException, ) +from aws_durable_execution_sdk_python_testing.execution import Execution +from aws_durable_execution_sdk_python_testing.model import StartDurableExecutionInput + +PARENT_ACCOUNT = "123456789012" + + +def _execution(region: str | None = "us-west-2") -> Execution: + execution = Execution.new( + StartDurableExecutionInput( + account_id=PARENT_ACCOUNT, + function_name="parent", + function_qualifier="$LATEST", + execution_name="run", + execution_timeout_seconds=60, + execution_retention_period_days=1, + ) + ) + execution.region = region + return execution -def test_validate_start_action_with_no_current_state(): - """Test START action with no current state.""" - update = OperationUpdate( - operation_id="test-id", +def _start_update( + operation_id: str = "test-id", + function_name: str = "child-function", + tenant_id: str | None = None, +) -> OperationUpdate: + return OperationUpdate( + operation_id=operation_id, operation_type=OperationType.CHAINED_INVOKE, action=OperationAction.START, + chained_invoke_options=ChainedInvokeOptions( + function_name=function_name, tenant_id=tenant_id + ), + ) + + +def test_validate_start_action_with_no_current_state(): + """Test START action with no current state.""" + ChainedInvokeOperationValidator.validate(None, _start_update(), _execution()) + + +@pytest.mark.parametrize( + "function_name", + [ + "child", + "child:prod", + f"{PARENT_ACCOUNT}:function:child", + f"arn:aws:lambda:us-west-2:{PARENT_ACCOUNT}:function:child", + f"arn:aws:lambda:us-west-2:{PARENT_ACCOUNT}:function:child:$LATEST", + "namespace.child", + "child:$LATEST.PUBLISHED", + "x" * 256, + ], +) +def test_validate_accepts_every_lambda_function_name_form(function_name: str): + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=function_name), _execution() + ) + + +@pytest.mark.parametrize( + "function_name", + [ + "", + "has space", + "arn:aws:lambda:us-west-2:function:child", + "arn:aws:lambda:function:child", + "x" * 257, + "child\n", + ], +) +def test_validate_rejects_malformed_target(function_name: str): + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=function_name), _execution() + ) + assert str(exc_info.value) == f"Invalid function ARN '{function_name}'" + + +def test_validate_rejects_target_in_another_account(): + target = "arn:aws:lambda:us-west-2:999999999999:function:child" + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=target), _execution() + ) + assert ( + str(exc_info.value) + == "Cannot start a CHAINED_INVOKE on a function in another account." + ) + + +def test_validate_rejects_target_in_another_region(): + target = f"arn:aws:lambda:eu-west-1:{PARENT_ACCOUNT}:function:child" + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=target), _execution() + ) + assert ( + str(exc_info.value) + == "Cannot start a CHAINED_INVOKE on a function in another region." + ) + + +@pytest.mark.parametrize("tenant_id", ["tenant-a", "a.b_c:d/e=f+g-h@i j", "x" * 256]) +def test_validate_accepts_valid_tenant_ids(tenant_id: str): + ChainedInvokeOperationValidator.validate( + None, _start_update(tenant_id=tenant_id), _execution() + ) + + +@pytest.mark.parametrize( + "tenant_id", ["", "x" * 257, "bad#tenant", "tab\there", "tenant\n"] +) +def test_validate_rejects_invalid_tenant_ids(tenant_id: str): + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(tenant_id=tenant_id), _execution() + ) + assert str(exc_info.value) == ( + "TenantId must be 1 to 256 characters matching [a-zA-Z0-9._:/=+-@ ]." + ) + + +def test_validate_skips_region_check_when_execution_has_no_region(): + target = f"arn:aws:lambda:eu-west-1:{PARENT_ACCOUNT}:function:child" + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=target), _execution(region=None) ) - ChainedInvokeOperationValidator.validate(None, update) def test_validate_start_action_with_existing_state(): @@ -34,76 +153,77 @@ def test_validate_start_action_with_existing_state(): operation_type=OperationType.CHAINED_INVOKE, status=OperationStatus.STARTED, ) - update = OperationUpdate( - operation_id="test-id", - operation_type=OperationType.CHAINED_INVOKE, - action=OperationAction.START, - ) with pytest.raises( InvalidParameterValueException, - match="Cannot start an INVOKE that already exist", + match="Cannot start a CHAINED_INVOKE operation that already exists", ): - ChainedInvokeOperationValidator.validate(current_state, update) - - -def test_validate_cancel_action_with_started_state(): - """Test CANCEL action with STARTED state.""" - current_state = Operation( - operation_id="test-id", - operation_type=OperationType.CHAINED_INVOKE, - status=OperationStatus.STARTED, - ) - update = OperationUpdate( - operation_id="test-id", - operation_type=OperationType.CHAINED_INVOKE, - action=OperationAction.CANCEL, - ) - ChainedInvokeOperationValidator.validate(current_state, update) + ChainedInvokeOperationValidator.validate( + current_state, _start_update(), _execution() + ) -def test_validate_cancel_action_with_no_current_state(): - """Test CANCEL action with no current state raises error.""" +def test_validate_start_action_without_options(): + """Test START action without ChainedInvokeOptions raises error.""" update = OperationUpdate( operation_id="test-id", operation_type=OperationType.CHAINED_INVOKE, - action=OperationAction.CANCEL, + action=OperationAction.START, ) with pytest.raises( InvalidParameterValueException, - match="Cannot cancel an INVOKE that does not exist or has already completed", + match="Update for CHAINED_INVOKE operation requires ChainedInvokeOptions", ): - ChainedInvokeOperationValidator.validate(None, update) + ChainedInvokeOperationValidator.validate(None, update, _execution()) -def test_validate_cancel_action_with_completed_state(): - """Test CANCEL action with completed state raises error.""" +@pytest.mark.parametrize( + "action", + [ + OperationAction.CANCEL, + OperationAction.SUCCEED, + OperationAction.FAIL, + OperationAction.RETRY, + ], +) +def test_validate_non_start_action_rejected(action: OperationAction): + """Test every non-START action raises error.""" current_state = Operation( operation_id="test-id", operation_type=OperationType.CHAINED_INVOKE, - status=OperationStatus.SUCCEEDED, + status=OperationStatus.STARTED, ) update = OperationUpdate( operation_id="test-id", operation_type=OperationType.CHAINED_INVOKE, - action=OperationAction.CANCEL, + action=action, ) with pytest.raises( InvalidParameterValueException, - match="Cannot cancel an INVOKE that does not exist or has already completed", + match="Invalid action for CHAINED_INVOKE operation", ): - ChainedInvokeOperationValidator.validate(current_state, update) - - -def test_validate_invalid_action(): - """Test invalid action raises error.""" - update = OperationUpdate( - operation_id="test-id", - operation_type=OperationType.CHAINED_INVOKE, - action=OperationAction.SUCCEED, + ChainedInvokeOperationValidator.validate(current_state, update, _execution()) + + +@pytest.mark.parametrize("function_name", [None, 42, ["child"]]) +def test_validate_rejects_a_non_string_function_name(function_name): + """The API model requires FunctionName as a string. Another type is a + 400 validation error, not a 500 from a TypeError.""" + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(function_name=function_name), _execution() + ) + assert str(exc_info.value) == f"Invalid function ARN '{function_name}'" + + +@pytest.mark.parametrize("tenant_id", [7, 1.5, ["t"]]) +def test_validate_rejects_a_non_string_tenant_id(tenant_id): + with pytest.raises(InvalidParameterValueException) as exc_info: + ChainedInvokeOperationValidator.validate( + None, _start_update(tenant_id=tenant_id), _execution() + ) + assert str(exc_info.value) == ( + "TenantId must be 1 to 256 characters matching [a-zA-Z0-9._:/=+-@ ]." ) - - with pytest.raises(InvalidParameterValueException, match="Invalid INVOKE action"): - ChainedInvokeOperationValidator.validate(None, update) diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/valid_actions_by_operation_type_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/valid_actions_by_operation_type_test.py index 1fce1938d..1ce9652c5 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/valid_actions_by_operation_type_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/checkpoint/validators/valid_actions_by_operation_type_test.py @@ -64,7 +64,6 @@ def test_validate_invoke_valid_actions(): """Test valid actions for INVOKE operations.""" valid_actions = [ OperationAction.START, - OperationAction.CANCEL, ] for action in valid_actions: CheckpointValidator._validate_valid_action_for_type( diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/child_dispatcher_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/child_dispatcher_test.py new file mode 100644 index 000000000..01e3c1077 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/child_dispatcher_test.py @@ -0,0 +1,709 @@ +"""Unit tests for chained-invoke dispatchers.""" + +import io +import json +import re +from dataclasses import replace +import threading +from types import SimpleNamespace + +import pytest +from typing import Any + +from botocore.exceptions import ClientError, ReadTimeoutError # type: ignore + +from aws_durable_execution_sdk_python_testing.clock import RealClock +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + CHAINED_INVOKE_TIMEOUT_ERROR_TYPE, + DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS, + DEFAULT_CHILD_RETENTION_PERIOD_DAYS, + FUNCTION_NOT_FOUND_ERROR_TYPE, + FUNCTION_TIMEOUT_ERROR_TYPE, + MAX_CHAINED_INVOKE_PAYLOAD_BYTES, + ChainedInvokeRequest, + ChildOutcome, + EndpointChildDispatcher, + FunctionConfig, + FunctionConfigs, + FunctionRegistry, + InProcessChildDispatcher, + KnownOutcome, + RunInvocation, + StartChild, + UnconfiguredChildDispatcher, + invoke_function, + is_valid_tenant_id, + parse_function_target, +) + + +TIMEOUT = 5 # invocation timeout used by the dispatchers under test + + +def _request(function_name: str = "target", payload: str | None = None): + return ChainedInvokeRequest( + parent_execution_arn="parent-arn", + operation_id="op-1", + function_name=function_name, + tenant_id=None, + payload=payload, + account_id="123456789012", + ) + + +def _fake_context(request=None) -> object: # noqa: ARG001 + return SimpleNamespace(aws_request_id="req-1") + + +# region InProcessChildDispatcher + + +def test_in_process_unknown_function_fails_to_start(): + dispatcher = InProcessChildDispatcher( + FunctionRegistry(), _fake_context, TIMEOUT, RealClock() + ) + result = dispatcher.dispatch(_request("missing")) + assert isinstance(result, KnownOutcome) + assert result.outcome.error is not None + assert result.outcome.error.type == FUNCTION_NOT_FOUND_ERROR_TYPE + assert result.outcome.error.message is not None + assert result.outcome.error.message.startswith("Function not found: missing.") + + +def test_in_process_preflight_resolves_registrations_without_invoking(): + """The preflight answers from the registry alone: a registered name + (qualified or bare) passes, an unknown one fails as not found.""" + registry = FunctionRegistry() + registry.register("child", lambda event, context: None, is_durable=False) + registry.register("exact:prod", lambda event, context: None, is_durable=True) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + + assert dispatcher.preflight(parse_function_target("child")) is None + assert dispatcher.preflight(parse_function_target("child:staging")) is None + assert dispatcher.preflight(parse_function_target("exact:prod")) is None + + outcome = dispatcher.preflight(parse_function_target("missing")) + assert outcome is not None + assert outcome.error is not None + assert outcome.error.type == FUNCTION_NOT_FOUND_ERROR_TYPE + assert outcome.error.message is not None + assert outcome.error.message.startswith("Function not found: missing.") + + +def test_in_process_durable_function_returns_child_start(): + registry = FunctionRegistry() + registry.register( + "target", + lambda event, context: None, + is_durable=True, + execution_timeout_seconds=42, + retention_period_days=3, + ) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + result = dispatcher.dispatch(_request(payload='{"n": 1}')) + assert isinstance(result, StartChild) + assert result.child_start.function_name == "target" + assert result.child_start.execution_timeout_seconds == 42 + assert result.child_start.execution_retention_period_days == 3 + assert result.child_start.input == '{"n": 1}' + + +def test_in_process_non_durable_function_runs_once(): + registry = FunctionRegistry() + registry.register( + "target", + lambda event, context: {"total": event["a"] + event["b"]}, + is_durable=False, + ) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + result = dispatcher.dispatch(_request(payload='{"a": 2, "b": 3}')) + assert isinstance(result, RunInvocation) + outcome = result.invocation() + assert isinstance(outcome, ChildOutcome) + assert json.loads(outcome.result) == {"total": 5} + + +def test_in_process_non_durable_function_error_becomes_outcome(): + def exploding(event: Any, context: Any) -> None: + msg: str = "boom" + raise ValueError(msg) + + registry = FunctionRegistry() + registry.register("target", exploding, is_durable=False) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + dispatched = dispatcher.dispatch(_request()) + assert isinstance(dispatched, RunInvocation) + outcome = dispatched.invocation() + assert isinstance(outcome, ChildOutcome) + assert outcome.error is not None + assert "boom" in (outcome.error.message or "") + + +def test_in_process_non_durable_function_exceeding_timeout_fails_as_lambda_does(): + """A plain target running past the invocation timeout ends FAILED with + Lambda's function-timeout error, as the service reports it.""" + release = threading.Event() + + def slow(event, context): # noqa: ARG001 + release.wait(5) + return "late" + + registry = FunctionRegistry() + registry.register("target", slow, is_durable=False) + dispatcher = InProcessChildDispatcher(registry, _fake_context, 1, RealClock()) + dispatched = dispatcher.dispatch(_request()) + assert isinstance(dispatched, RunInvocation) + try: + outcome = dispatched.invocation() + finally: + release.set() + assert outcome.timed_out is False + assert outcome.result is None + assert outcome.error is not None + assert outcome.error.type == FUNCTION_TIMEOUT_ERROR_TYPE == "Sandbox.Timedout" + assert re.fullmatch( + r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}\.\d{3}Z req-1 Task timed out after 1\.00 seconds", + outcome.error.message or "", + ) + + +def test_in_process_non_durable_result_over_limit_fails(): + registry = FunctionRegistry() + registry.register( + "target", + lambda event, context: "x" * MAX_CHAINED_INVOKE_PAYLOAD_BYTES, # noqa: ARG005 + is_durable=False, + ) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + dispatched = dispatcher.dispatch(_request()) + assert isinstance(dispatched, RunInvocation) + outcome = dispatched.invocation() + assert outcome.error is not None + assert outcome.error.type is None + assert outcome.error.message == ( + "CHAINED_INVOKE output payload size must be less than or equal to " + "1048576 bytes." + ) + + +# endregion + +# region invoke_function (non-durable targets at a Lambda endpoint) + + +class _StubLambdaClient: + def __init__(self, response: dict | None = None, error: Exception | None = None): + self.response = response + self.error = error + self.calls: list[dict] = [] + + def invoke(self, **kwargs): + self.calls.append(kwargs) + if self.error is not None: + raise self.error + return self.response + + +def _response(body: str = "", function_error: bool = False) -> dict: + response: dict = { + "StatusCode": 200, + "Payload": io.BytesIO(body.encode()), + "ResponseMetadata": {"HTTPHeaders": {}}, + } + if function_error: + response["FunctionError"] = "Unhandled" + return response + + +def test_invoke_function_returns_body_as_result(): + client = _StubLambdaClient(response=_response(body='{"total": 5}')) + result = invoke_function(client, "target", '{"n": 1}', None, TIMEOUT) + assert result == ChildOutcome(result='{"total": 5}') + assert client.calls[0]["InvocationType"] == "RequestResponse" + assert client.calls[0]["FunctionName"] == "target" + assert client.calls[0]["Payload"] == '{"n": 1}' + + +def test_invoke_function_passes_tenant_id_when_present(): + client = _StubLambdaClient(response=_response(body="1")) + invoke_function(client, "target", "{}", "tenant-a", TIMEOUT) + assert client.calls[0]["TenantId"] == "tenant-a" + + +def test_invoke_function_absent_payload_delivers_empty_object(): + client = _StubLambdaClient(response=_response(body="1")) + invoke_function(client, "target", None, None, TIMEOUT) + assert client.calls[0]["Payload"] == "{}" + assert "TenantId" not in client.calls[0] + + +def test_invoke_function_empty_body_is_no_result(): + client = _StubLambdaClient(response=_response(body="")) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result == ChildOutcome(result=None) + + +def test_invoke_function_error_maps_to_error_outcome(): + body = json.dumps( + {"errorMessage": "boom", "errorType": "ValueError", "stackTrace": []} + ) + client = _StubLambdaClient(response=_response(body=body, function_error=True)) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert result.error.message == "boom" + assert result.error.type == "ValueError" + + +def test_invoke_function_error_with_non_json_body_keeps_body_as_message(): + client = _StubLambdaClient(response=_response(body="not json", function_error=True)) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert result.error.message == "not json" + + +def test_invoke_function_api_error_surfaces_lambda_error_code(): + client = _StubLambdaClient( + error=ClientError( + { + "Error": { + "Code": "ResourceNotFoundException", + "Message": "Function not found: arn:aws:lambda:x:1:function:target", + } + }, + "Invoke", + ) + ) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert result.error.type == "ResourceNotFoundException" + assert ( + result.error.message == "Function not found: arn:aws:lambda:x:1:function:target" + ) + + +def test_invoke_function_transport_error_fails_without_error_type(): + client = _StubLambdaClient(error=ConnectionError("refused")) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert result.error.type is None + assert result.error.message == "Failed to invoke target: refused" + + +def test_invoke_function_read_timeout_is_a_chained_invoke_timeout(): + """The endpoint held the Invoke past the invocation timeout.""" + client = _StubLambdaClient( + error=ReadTimeoutError(endpoint_url="http://localhost:3001") + ) + result = invoke_function(client, "target", None, None, 900) + assert result.timed_out is True + assert result.error is not None + assert result.error.type == CHAINED_INVOKE_TIMEOUT_ERROR_TYPE + assert result.error.message == "CHAINED_INVOKE timed out after 900 seconds." + + +def test_invoke_function_body_read_timeout_is_a_chained_invoke_timeout(): + """The endpoint answered with headers, then held the body past the + invocation timeout. The body read times out on its own and is the + same chained-invoke timeout, not a generic failure.""" + + class _TimingOutBody: + def read(self) -> bytes: + raise ReadTimeoutError(endpoint_url="http://localhost:3001") + + client = _StubLambdaClient(response={"Payload": _TimingOutBody()}) + result = invoke_function(client, "target", None, None, 900) + assert result.timed_out is True + assert result.error is not None + assert result.error.type == CHAINED_INVOKE_TIMEOUT_ERROR_TYPE + assert result.error.message == "CHAINED_INVOKE timed out after 900 seconds." + + +def test_invoke_function_result_over_limit_fails(): + client = _StubLambdaClient( + response=_response(body="x" * (MAX_CHAINED_INVOKE_PAYLOAD_BYTES + 1)) + ) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.result is None + assert result.error is not None + assert result.error.message == ( + "CHAINED_INVOKE output payload size must be less than or equal to " + "1048576 bytes." + ) + + +def test_invoke_function_result_at_limit_is_accepted(): + body = "x" * MAX_CHAINED_INVOKE_PAYLOAD_BYTES + client = _StubLambdaClient(response=_response(body=body)) + assert invoke_function(client, "target", None, None, TIMEOUT).result == body + + +def test_invoke_function_non_utf8_body_fails_with_service_message(): + response = _response() + response["Payload"] = io.BytesIO(b"\xff\xfe") + client = _StubLambdaClient(response=response) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert ( + result.error.message == "CHAINED_INVOKE response could not be decoded as UTF-8." + ) + + +def test_invoke_function_error_keeps_error_data(): + body = json.dumps( + { + "errorMessage": "boom", + "errorType": "ValueError", + "errorData": '{"code": 7}', + "stackTrace": ["frame"], + } + ) + client = _StubLambdaClient(response=_response(body=body, function_error=True)) + result = invoke_function(client, "target", None, None, TIMEOUT) + assert result.error is not None + assert result.error.data == '{"code": 7}' + assert result.error.stack_trace == ["frame"] + + +# endregion + +# region parse_function_target + + +def test_parse_function_target_forms(): + assert parse_function_target("child").name == "child" + assert parse_function_target("child:prod").qualifier == "prod" + partial = parse_function_target("123456789012:function:child") + assert (partial.name, partial.account_id, partial.region) == ( + "child", + "123456789012", + None, + ) + full = parse_function_target( + "arn:aws:lambda:us-west-2:123456789012:function:child:$LATEST" + ) + assert (full.name, full.qualifier, full.account_id, full.region) == ( + "child", + "$LATEST", + "123456789012", + "us-west-2", + ) + + +def test_parse_function_target_accepts_invoke_forms(): + """Dotted names, the $LATEST.PUBLISHED qualifier, and 256 characters.""" + assert parse_function_target("namespace.child").name == "namespace.child" + assert parse_function_target("child:$LATEST.PUBLISHED").qualifier == ( + "$LATEST.PUBLISHED" + ) + assert parse_function_target("x" * 256).name == "x" * 256 + + +def test_parse_function_target_rejects_malformed(): + for bad in ( + "", + "a b", + "arn:aws:lambda:us-west-2:function:child", + "x" * 257, + "child\n", + "child:prod\n", + ): + with pytest.raises(ValueError, match=".*"): + parse_function_target(bad) + + +def test_is_valid_tenant_id_requires_the_whole_value_to_match(): + assert is_valid_tenant_id("tenant-a") + assert not is_valid_tenant_id("tenant\n") + assert not is_valid_tenant_id("") + + +# endregion + +# region qualifier resolution + + +def test_lookup_keys_prefer_the_qualified_registration(): + request = _request() + assert request.lookup_keys() == ("target",) + assert request.child_qualifier() == "$LATEST" + qualified = ChainedInvokeRequest( + parent_execution_arn="parent-arn", + operation_id="op-1", + function_name="target", + qualifier="prod", + tenant_id=None, + payload=None, + account_id="123456789012", + ) + assert qualified.lookup_keys() == ("target:prod", "target") + assert qualified.child_qualifier() == "prod" + assert qualified.invoke_identifier() == "target:prod" + assert request.invoke_identifier() == "target" + assert replace(qualified, qualifier="$LATEST").invoke_identifier() == ( + "target:$LATEST" + ) + + +def _qualified_request(qualifier: str) -> ChainedInvokeRequest: + return ChainedInvokeRequest( + parent_execution_arn="parent-arn", + operation_id="op-1", + function_name="target", + qualifier=qualifier, + tenant_id=None, + payload="{}", + account_id="123456789012", + ) + + +def test_in_process_qualified_registration_selects_the_registration_only(): + """The child keeps the requested name and qualifier whichever key matched.""" + registry = FunctionRegistry() + registry.register("target", lambda e, c: "latest", is_durable=True) # noqa: ARG005 + registry.register( + "target:prod", + lambda e, c: "prod", # noqa: ARG005 + is_durable=True, + execution_timeout_seconds=7, + ) + dispatcher = InProcessChildDispatcher(registry, _fake_context, TIMEOUT, RealClock()) + + result = dispatcher.dispatch(_qualified_request("prod")) + assert isinstance(result, StartChild) + assert result.child_start.function_name == "target" + assert result.child_start.function_qualifier == "prod" + assert result.child_start.execution_timeout_seconds == 7 + + fallback = dispatcher.dispatch(_qualified_request("3")) + assert isinstance(fallback, StartChild) + assert fallback.child_start.function_name == "target" + assert fallback.child_start.function_qualifier == "3" + assert fallback.child_start.execution_timeout_seconds != 7 + + +def test_endpoint_invokes_the_requested_identifier_whichever_config_matched(): + client = _StubLambdaClient(response=_response(body="1")) + dispatcher = EndpointChildDispatcher( + FunctionConfigs({"target": FunctionConfig(is_durable=False)}), + lambda _arn: client, + TIMEOUT, + ) + + dispatched = dispatcher.dispatch(_qualified_request("staging")) + assert isinstance(dispatched, RunInvocation) + dispatched.invocation() + assert client.calls[-1]["FunctionName"] == "target:staging" + + dispatched = dispatcher.dispatch(_request()) + assert isinstance(dispatched, RunInvocation) + dispatched.invocation() + assert client.calls[-1]["FunctionName"] == "target" + + +# endregion + +# region UnconfiguredChildDispatcher + + +def test_unconfigured_dispatcher_fails_every_invoke_naming_the_option(): + result = UnconfiguredChildDispatcher().dispatch(_request("child")) + assert isinstance(result, KnownOutcome) + assert result.outcome.error is not None + assert result.outcome.error.type is None + assert result.outcome.error.message is not None + assert result.outcome.error.message.startswith("Cannot invoke child:") + assert "--function-configs" in result.outcome.error.message + + +def test_unconfigured_dispatcher_fails_the_preflight_of_every_target(): + outcome = UnconfiguredChildDispatcher().preflight(parse_function_target("child")) + assert outcome is not None + assert outcome.error is not None + assert outcome.error.message is not None + assert outcome.error.message.startswith("Cannot invoke child:") + assert "--function-configs" in outcome.error.message + + +# endregion + +# region EndpointChildDispatcher + + +def test_function_config_without_durable_config_is_plain(): + for entry in ({}, None, {"Timeout": 30}): + config = FunctionConfig.from_dict(entry) + assert config.is_durable is False + + +def test_function_config_with_empty_durable_config_is_durable_with_defaults(): + config = FunctionConfig.from_dict({"DurableConfig": {}}) + assert config.is_durable is True + assert config.execution_timeout_seconds == DEFAULT_CHILD_EXECUTION_TIMEOUT_SECONDS + assert config.retention_period_days == DEFAULT_CHILD_RETENTION_PERIOD_DAYS + + +def test_function_config_reads_durable_config_fields(): + config = FunctionConfig.from_dict( + {"DurableConfig": {"ExecutionTimeout": 45, "RetentionPeriodInDays": 3}} + ) + assert config.is_durable is True + assert config.execution_timeout_seconds == 45 + assert config.retention_period_days == 3 + + +@pytest.mark.parametrize( + ("entry", "detail"), + [ + ([], "must be a JSON object or null"), + (False, "must be a JSON object or null"), + (0, "must be a JSON object or null"), + ("", "must be a JSON object or null"), + ("durable", "must be a JSON object or null"), + ({"DurableConfig": True}, "DurableConfig must be a JSON object"), + ({"DurableConfig": []}, "DurableConfig must be a JSON object"), + ({"DurableConfig": "yes"}, "DurableConfig must be a JSON object"), + ({"DurableConfig": {"ExecutionTimeout": "60"}}, "ExecutionTimeout"), + ({"DurableConfig": {"ExecutionTimeout": 0}}, "ExecutionTimeout"), + ({"DurableConfig": {"ExecutionTimeout": True}}, "ExecutionTimeout"), + ({"DurableConfig": {"RetentionPeriodInDays": 1.5}}, "RetentionPeriodInDays"), + ], +) +def test_function_config_rejects_a_malformed_entry(entry, detail): + """A typo must not turn a durable target plain, nor crash with a + bare traceback: every wrong shape is a ValueError naming the function.""" + with pytest.raises(ValueError, match=detail) as exc_info: + FunctionConfig.from_dict(entry, "process-payment") + assert "'process-payment'" in str(exc_info.value) + + +def test_function_configs_name_the_malformed_function(): + with pytest.raises(ValueError, match="'lookup-price'.*JSON object or null"): + FunctionConfigs.from_value('{"process-payment": {}, "lookup-price": []}') + + +def test_function_configs_parse_inline_json(): + configs = FunctionConfigs.from_value( + '{"process-payment": {"DurableConfig": {"ExecutionTimeout": 60}}, "lookup-price": {}}' + ) + assert configs.by_name["process-payment"] == FunctionConfig( + is_durable=True, execution_timeout_seconds=60 + ) + assert configs.by_name["lookup-price"].is_durable is False + + +def test_function_configs_reject_a_non_object(): + with pytest.raises(ValueError, match="JSON object"): + FunctionConfigs.from_value('["process-payment"]') + + +def test_function_configs_read_a_file_url(tmp_path): + path = tmp_path / "function-configs.json" + path.write_text( + json.dumps( + { + "process-payment": {"DurableConfig": {"ExecutionTimeout": 60}}, + "lookup-price": {}, + "null_entry": None, + } + ), + encoding="utf-8", + ) + configs = FunctionConfigs.from_value(f"file://{path}") + assert set(configs.by_name) == {"process-payment", "lookup-price", "null_entry"} + assert configs.by_name["process-payment"] == FunctionConfig( + is_durable=True, execution_timeout_seconds=60 + ) + assert configs.by_name["lookup-price"].is_durable is False + assert configs.by_name["null_entry"].is_durable is False + + +def test_endpoint_preflight_resolves_configurations_without_the_endpoint(): + class _NoClient: + def __call__(self, _arn): + msg = "the preflight must not touch the endpoint" + raise AssertionError(msg) + + configs = FunctionConfigs.from_value( + '{"child": {}, "exact:prod": {"DurableConfig": {}}}' + ) + dispatcher = EndpointChildDispatcher(configs, _NoClient(), TIMEOUT) + + assert dispatcher.preflight(parse_function_target("child")) is None + assert dispatcher.preflight(parse_function_target("child:v2")) is None + assert dispatcher.preflight(parse_function_target("exact:prod")) is None + + outcome = dispatcher.preflight(parse_function_target("missing")) + assert outcome is not None + assert outcome.error is not None + assert outcome.error.type == FUNCTION_NOT_FOUND_ERROR_TYPE + assert outcome.error.message == "Function not found: missing." + + +def test_endpoint_unknown_function_fails_to_start(): + dispatcher = EndpointChildDispatcher( + FunctionConfigs({}), lambda _arn: _StubLambdaClient(), TIMEOUT + ) + result = dispatcher.dispatch(_request("missing")) + assert isinstance(result, KnownOutcome) + assert result.outcome.error is not None + assert result.outcome.error.type == FUNCTION_NOT_FOUND_ERROR_TYPE + assert result.outcome.error.message == "Function not found: missing." + + +def test_endpoint_durable_function_returns_child_start_from_config(): + configs = { + "target": FunctionConfig( + is_durable=True, execution_timeout_seconds=30, retention_period_days=2 + ) + } + dispatcher = EndpointChildDispatcher( + FunctionConfigs(configs), lambda _arn: _StubLambdaClient(), TIMEOUT + ) + result = dispatcher.dispatch(_request(payload='{"n": 1}')) + assert isinstance(result, StartChild) + start = result.child_start + assert start.function_name == "target" + assert start.account_id == "123456789012" + assert start.execution_timeout_seconds == 30 + assert start.execution_retention_period_days == 2 + assert start.input == '{"n": 1}' + assert start.lambda_endpoint is None + + +def test_endpoint_non_durable_function_uses_the_parent_execution_client(): + """The client is asked for per dispatch, for the parent execution, so a + plain target goes where that execution's own invocations go.""" + clients = { + "arn:parent-a": _StubLambdaClient(response=_response(body="1")), + "arn:parent-b": _StubLambdaClient(response=_response(body="2")), + } + dispatcher = EndpointChildDispatcher( + FunctionConfigs({"target": FunctionConfig(is_durable=False)}), + clients.__getitem__, + TIMEOUT, + ) + + for arn, expected in (("arn:parent-a", "1"), ("arn:parent-b", "2")): + dispatched = dispatcher.dispatch(replace(_request(), parent_execution_arn=arn)) + assert isinstance(dispatched, RunInvocation) + assert dispatched.invocation().result == expected + assert len(clients["arn:parent-a"].calls) == 1 + assert len(clients["arn:parent-b"].calls) == 1 + + +def test_endpoint_non_durable_function_invokes_at_endpoint(): + client = _StubLambdaClient(response=_response(body='{"total": 5}')) + configs = {"target": FunctionConfig(is_durable=False)} + dispatcher = EndpointChildDispatcher( + FunctionConfigs(configs), lambda _arn: client, TIMEOUT + ) + dispatched = dispatcher.dispatch(_request(payload='{"n": 1}')) + assert isinstance(dispatched, RunInvocation) + result = dispatched.invocation() + assert isinstance(result, ChildOutcome) + assert result.result == '{"total": 5}' + assert client.calls[0]["FunctionName"] == "target" + assert client.calls[0]["Payload"] == '{"n": 1}' + + +# endregion diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/cli_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/cli_test.py index 98b53c6ea..7b41f2611 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/cli_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/cli_test.py @@ -16,6 +16,7 @@ from botocore.exceptions import ConnectionError # type: ignore +from aws_durable_execution_sdk_python_testing.child_dispatcher import FunctionConfigs from aws_durable_execution_sdk_python_testing.cli import CliApp, CliConfig, main from aws_durable_execution_sdk_python_testing.exceptions import ( DurableFunctionsLocalRunnerError, @@ -195,6 +196,61 @@ def test_start_server_command_parses_arguments_correctly() -> None: assert exit_code == 130 # KeyboardInterrupt exit code +def test_start_server_parses_function_configs_at_the_boundary() -> None: + """--function-configs reaches the runner as FunctionConfigs, not text.""" + app = CliApp() + with patch( + "aws_durable_execution_sdk_python_testing.cli.WebRunner" + ) as mock_web_runner: + mock_runner_instance = mock_web_runner.return_value + mock_runner_instance.__enter__.return_value = mock_runner_instance + mock_runner_instance.__exit__.return_value = None + mock_runner_instance.serve_forever.side_effect = KeyboardInterrupt() + + app.run( + [ + "start-server", + "--function-configs", + '{"ProcessPayment": {"DurableConfig": {"ExecutionTimeout": 60}}, "LookupPrice": {}}', + ] + ) + + config = mock_web_runner.call_args.args[0] + assert isinstance(config.function_configs, FunctionConfigs) + assert config.function_configs.by_name["ProcessPayment"].is_durable + assert ( + config.function_configs.by_name["ProcessPayment"].execution_timeout_seconds + == 60 + ) + assert not config.function_configs.by_name["LookupPrice"].is_durable + + +@pytest.mark.parametrize( + ("value", "detail"), + [ + ('["ProcessPayment"]', "JSON object"), + ("{not json", "Expecting"), + ("file:///nonexistent/function-configs.json", "No such file"), + ('{"ProcessPayment": []}', "'ProcessPayment': the entry must be"), + ('{"ProcessPayment": false}', "'ProcessPayment': the entry must be"), + ('{"ProcessPayment": "plain"}', "'ProcessPayment': the entry must be"), + ( + '{"ProcessPayment": {"DurableConfig": true}}', + "'ProcessPayment': DurableConfig", + ), + ], +) +def test_start_server_rejects_a_bad_function_configs_value( + value: str, detail: str +) -> None: + """A malformed value is an argument error that names the option.""" + with patch("sys.stderr", new_callable=StringIO) as mock_stderr: + exit_code = CliApp().run(["start-server", "--function-configs", value]) + assert exit_code == 2 # argparse usage error + assert "--function-configs" in mock_stderr.getvalue() + assert detail in mock_stderr.getvalue() + + def test_invoke_command_parses_arguments_correctly() -> None: """Test that invoke command parses arguments correctly.""" app = CliApp() diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/e2e/chained_invoke_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/e2e/chained_invoke_test.py new file mode 100644 index 000000000..bb0dc7c6f --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/e2e/chained_invoke_test.py @@ -0,0 +1,664 @@ +"""End-to-end tests for chained invoke through the in-process runner.""" + +import json +import time +from typing import Any + +from aws_durable_execution_sdk_python.config import Duration, InvokeConfig +from aws_durable_execution_sdk_python.context import DurableContext +from aws_durable_execution_sdk_python.exceptions import InvokeError +from aws_durable_execution_sdk_python.execution import durable_execution +from aws_durable_execution_sdk_python.lambda_service import ( + ErrorObject, + OperationStatus, +) + +from aws_durable_execution_sdk_python_testing.runner import ( + DurableFunctionTestResult, + DurableFunctionTestRunner, +) + + +@durable_execution +def child_doubler(event: Any, context: DurableContext) -> int: + return event["n"] * 2 + + +@durable_execution +def child_failer(event: Any, context: DurableContext) -> int: + msg: str = "child exploded" + raise ValueError(msg) + + +def plain_adder(event: Any, context: Any) -> dict: + return {"total": event["a"] + event["b"]} + + +def test_invoke_durable_child_returns_result() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + doubled: int = context.invoke("child-doubler", {"n": 21}, name="double") + return doubled + 1 + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + result: DurableFunctionTestResult = runner.run(input=json.dumps({})) + + assert result.result == "43" + invoke_op = result.get_invoke("double") + assert invoke_op.status == OperationStatus.SUCCEEDED + assert invoke_op.result == "42" + + +def test_invoke_non_durable_child_returns_result() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + summed: dict = context.invoke("adder", {"a": 2, "b": 3}, name="add") + return summed["total"] + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_function("adder", plain_adder) + result = runner.run(input=json.dumps({})) + + assert result.result == "5" + assert result.get_invoke("add").status == OperationStatus.SUCCEEDED + + +def test_invoke_failing_durable_child_raises_in_parent() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("child-failer", {}, name="fail-target") + except InvokeError as err: + return f"caught:{err.message}" + return "not-reached" + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-failer", child_failer) + result = runner.run(input=json.dumps({})) + + assert "caught:" in (result.result or "") + assert result.get_invoke("fail-target").status == OperationStatus.FAILED + + +def test_invoke_unknown_function_fails_to_start() -> None: + """The service resolves a target before scheduling it and returns a + target it cannot resolve as FAILED in the checkpoint response, so the + handler raises without suspending. The runner does the same: the + parent finishes in its first invocation, and the history carries the + started and failed events of the operation.""" + invocations: list[int] = [] + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + invocations.append(1) + try: + context.invoke( + "nowhere-to-be-found", {"secret": "retain-me"}, name="missing" + ) + except InvokeError as err: + return f"type:{err.error_type}" + return "not-reached" + + with DurableFunctionTestRunner(handler=parent) as runner: + parent_arn: str = runner.run_async(input=json.dumps({})) + result = runner.wait_for_result(execution_arn=parent_arn, timeout=60) + history = runner.get_execution_history(parent_arn, include_execution_data=True) + + # An unknown target surfaces Lambda's error code for a missing function. + assert result.result == '"type:ResourceNotFoundException"' + assert result.get_invoke("missing").status == OperationStatus.FAILED + assert len(invocations) == 1 + chained_events = [ + e.event_type for e in history.events if e.event_type.startswith("ChainedInvoke") + ] + assert chained_events == ["ChainedInvokeStarted", "ChainedInvokeFailed"] + # The service keeps no input for a target it cannot resolve: the + # started event names the function and carries no Input. + started = next(e for e in history.events if e.event_type == "ChainedInvokeStarted") + assert started.chained_invoke_started_details is not None + assert started.chained_invoke_started_details.function_name == "nowhere-to-be-found" + assert started.chained_invoke_started_details.input is None + + +def test_invoke_stopped_child_surfaces_stop_error() -> None: + @durable_execution + def waiting_child(event: Any, context: DurableContext) -> str: + context.wait(Duration.from_hours(1), name="hold") + return "never" + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("waiting-child", {}, name="held") + except InvokeError as err: + return f"{err.error_type}:{err.message}" + return "not-reached" + + with DurableFunctionTestRunner(handler=parent, skip_time=False) as runner: + runner.register_durable_function("waiting-child", waiting_child) + parent_arn: str = runner.run_async(input=json.dumps({})) + + # Find the child through the parent's started event, then stop it + # with a caller-supplied error. + child_arn: str | None = None + deadline = time.monotonic() + 30 + while child_arn is None and time.monotonic() < deadline: + history = runner.get_execution_history(parent_arn) + for event in history.events: + details = event.chained_invoke_started_details + if details is not None and details.durable_execution_arn: + child_arn = details.durable_execution_arn + if child_arn is None: + time.sleep(0.05) + assert child_arn + + stop_error = ErrorObject( + message="operator stopped the child", + type="OperatorStop", + data=None, + stack_trace=None, + ) + runner._executor.stop_execution(child_arn, stop_error) # noqa: SLF001 + result = runner.wait_for_result(execution_arn=parent_arn, timeout=30) + + # The parent operation is STOPPED and carries the child's own stop error. + assert result.result == '"OperatorStop:operator stopped the child"' + held = result.get_invoke("held") + assert held.status == OperationStatus.STOPPED + assert held.error is not None + assert held.error.type == "OperatorStop" + assert held.error.message == "operator stopped the child" + + +def test_invoke_fan_out_multiple_children() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + first: int = context.invoke("child-doubler", {"n": 1}, name="one") + second: int = context.invoke("child-doubler", {"n": 2}, name="two") + third: dict = context.invoke("adder", {"a": first, "b": second}, name="three") + return third["total"] + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + runner.register_function("adder", plain_adder) + result = runner.run(input=json.dumps({})) + + assert result.result == "6" + for name in ("one", "two", "three"): + assert result.get_invoke(name).status == OperationStatus.SUCCEEDED + + +def test_invoke_nested_child_invokes_grandchild() -> None: + @durable_execution + def middle(event: Any, context: DurableContext) -> int: + doubled: int = context.invoke("child-doubler", {"n": event["n"]}, name="inner") + return doubled + 100 + + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + return context.invoke("middle", {"n": 5}, name="outer") + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("middle", middle) + runner.register_durable_function("child-doubler", child_doubler) + result = runner.run(input=json.dumps({})) + + assert result.result == "110" + + +def test_invoke_history_events_and_child_history() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + return context.invoke("child-doubler", {"n": 4}, name="double") + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + parent_arn: str = runner.run_async(input=json.dumps({})) + result = runner.wait_for_result(execution_arn=parent_arn, timeout=60) + + parent_history = runner.get_execution_history( + parent_arn, include_execution_data=True + ) + started_events = [ + e for e in parent_history.events if e.event_type == "ChainedInvokeStarted" + ] + assert len(started_events) == 1 + details = started_events[0].chained_invoke_started_details + assert details is not None + assert details.function_name == "child-doubler" + assert details.input is not None + assert json.loads(details.input.payload) == {"n": 4} + child_arn: str | None = details.durable_execution_arn + assert child_arn + + succeeded_events = [ + e for e in parent_history.events if e.event_type == "ChainedInvokeSucceeded" + ] + assert len(succeeded_events) == 1 + + # The child is a first-class execution with its own history. + child_history = runner.get_execution_history( + child_arn, include_execution_data=True + ) + child_event_types = [e.event_type for e in child_history.events] + assert "ExecutionStarted" in child_event_types + assert "ExecutionSucceeded" in child_event_types + + # Without execution data the started event redacts the payload. + redacted_history = runner.get_execution_history( + parent_arn, include_execution_data=False + ) + redacted_started = next( + e for e in redacted_history.events if e.event_type == "ChainedInvokeStarted" + ) + assert redacted_started.chained_invoke_started_details is not None + redacted_input = redacted_started.chained_invoke_started_details.input + assert redacted_input is not None + assert redacted_input.payload is None + assert redacted_input.truncated is True + + assert result.result == "8" + + +def test_invoke_durable_child_with_wait_under_skip_time() -> None: + @durable_execution + def slow_child(event: Any, context: DurableContext) -> str: + context.wait(Duration.from_hours(2), name="nap") + return "rested" + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + return context.invoke("slow-child", {}, name="patient") + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("slow-child", slow_child) + result = runner.run(input=json.dumps({}), execution_timeout=60) + + assert result.result == '"rested"' + + +def test_invoke_child_timeout_maps_to_timed_out() -> None: + @durable_execution + def hanging_child(event: Any, context: DurableContext) -> str: + time.sleep(2) + return "never" + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("hanging-child", {}, name="hang") + except InvokeError as err: + return f"type:{err.error_type}" + return "not-reached" + + with DurableFunctionTestRunner(handler=parent, skip_time=False) as runner: + runner.register_durable_function( + "hanging-child", hanging_child, execution_timeout=1 + ) + result = runner.run(input=json.dumps({}), execution_timeout=30) + + assert result.result == '"type:ChainedInvoke.Timeout"' + assert result.get_invoke("hang").status == OperationStatus.TIMED_OUT + + +def test_invoke_child_timeout_message_carries_the_child_execution_timeout() -> None: + @durable_execution + def hanging_child(event: Any, context: DurableContext) -> str: + time.sleep(2) + return "never" + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("hanging-child", {}, name="hang") + except InvokeError as err: + return str(err) + return "not-reached" + + with DurableFunctionTestRunner(handler=parent, skip_time=False) as runner: + runner.register_durable_function( + "hanging-child", hanging_child, execution_timeout=1 + ) + result = runner.run(input=json.dumps({}), execution_timeout=30) + + operation = result.get_invoke("hang") + assert operation.status == OperationStatus.TIMED_OUT + assert operation.error is not None + assert operation.error.type == "ChainedInvoke.Timeout" + assert operation.error.message == "CHAINED_INVOKE timed out after 1 seconds." + + +def test_invoke_plain_target_exceeding_invocation_timeout_fails_as_lambda_does() -> ( + None +): + """A non-durable target running past the invocation timeout fails with + Lambda's function-timeout error, as the service reports it.""" + + def slow_adder(event: Any, context: Any) -> dict: + time.sleep(3) + return {"total": 0} + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("slow-adder", {}, name="slow") + except InvokeError as err: + return f"type:{err.error_type}" + return "not-reached" + + with DurableFunctionTestRunner( + handler=parent, skip_time=False, invocation_timeout=1 + ) as runner: + runner.register_function("slow-adder", slow_adder) + result = runner.run(input=json.dumps({}), execution_timeout=30) + + assert result.result == '"type:Sandbox.Timedout"' + operation = result.get_invoke("slow") + assert operation.status == OperationStatus.FAILED + assert operation.error is not None + assert operation.error.message is not None + assert operation.error.message.endswith(" Task timed out after 1.00 seconds") + + +def test_invoke_target_given_as_arn_resolves_the_registered_name() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> int: + return context.invoke( + "arn:aws:lambda:us-west-2:123456789012:function:child-doubler", + {"n": 4}, + name="double", + ) + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + result = runner.run(input=json.dumps({})) + + assert result.result == "8" + assert result.get_invoke("double").status == OperationStatus.SUCCEEDED + + +def test_invoke_target_in_another_account_fails_the_checkpoint() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + context.invoke( + "arn:aws:lambda:us-west-2:999999999999:function:child-doubler", + {"n": 4}, + name="double", + ) + return "not-reached" + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + result = runner.run(input=json.dumps({})) + + assert result.error is not None + assert "Cannot start a CHAINED_INVOKE on a function in another account." in ( + result.error.message or "" + ) + + +def test_invoke_input_over_one_mebibyte_fails_the_checkpoint() -> None: + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + context.invoke("child-doubler", {"blob": "x" * 1_048_577}, name="big") + return "not-reached" + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child-doubler", child_doubler) + result = runner.run(input=json.dumps({})) + + assert result.error is not None + assert ( + "CHAINED_INVOKE input payload size must be less than or equal to " + "1048576 bytes." in (result.error.message or "") + ) + + +def test_invoke_durable_child_result_over_one_mebibyte_fails_the_child() -> None: + @durable_execution + def big_child(event: Any, context: DurableContext) -> str: + return "x" * 1_048_577 + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + try: + context.invoke("big-child", {}, name="big") + except InvokeError as err: + return str(err) + return "not-reached" + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("big-child", big_child) + result = runner.run(input=json.dumps({})) + + operation = result.get_invoke("big") + assert operation.status == OperationStatus.FAILED + assert operation.error is not None + assert ( + "Execution output payload size must be less than or equal to 1048576 bytes." + in (operation.error.message or "") + ) + + +def test_invoke_tenant_reaches_durable_and_plain_target_handlers() -> None: + """A tenant given on the invoke is what the target handlers see.""" + + @durable_execution + def tenant_echo(event: Any, context: DurableContext) -> str | None: + assert context.lambda_context is not None + return context.lambda_context.tenant_id + + def plain_tenant_echo(event: Any, context: Any) -> str | None: + return context.tenant_id + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + durable_seen: Any = context.invoke( + "tenant-echo", {}, name="durable", config=InvokeConfig(tenant_id="tenant-a") + ) + plain_seen: Any = context.invoke( + "plain-echo", {}, name="plain", config=InvokeConfig(tenant_id="tenant-b") + ) + return {"durable": durable_seen, "plain": plain_seen} + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("tenant-echo", tenant_echo) + runner.register_function("plain-echo", plain_tenant_echo) + result = runner.run(input=json.dumps({})) + + assert json.loads(result.result or "{}") == { + "durable": "tenant-a", + "plain": "tenant-b", + } + + +def test_invoke_region_of_the_runner_is_reported_and_enforced() -> None: + @durable_execution + def arn_echo(event: Any, context: DurableContext) -> str: + assert context.lambda_context is not None + return str(context.lambda_context.invoked_function_arn) + + @durable_execution + def parent(event: Any, context: DurableContext) -> str: + return context.invoke( + "arn:aws:lambda:eu-west-1:123456789012:function:arn-echo", + {}, + name="echo", + ) + + with DurableFunctionTestRunner(handler=parent, region="eu-west-1") as runner: + runner.register_durable_function("arn-echo", arn_echo) + result = runner.run(input=json.dumps({})) + + assert result.result == '"arn:aws:lambda:eu-west-1:123456789012:function:arn-echo"' + + +def test_invoke_without_tenant_gives_handlers_no_tenant() -> None: + @durable_execution + def tenant_echo(event: Any, context: DurableContext) -> str | None: + assert context.lambda_context is not None + return context.lambda_context.tenant_id + + def plain_tenant_echo(event: Any, context: Any) -> str | None: + return context.tenant_id + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + assert context.lambda_context is not None + durable_seen: Any = context.invoke("tenant-echo", {}, name="durable") + plain_seen: Any = context.invoke("plain-echo", {}, name="plain") + return { + "parent": context.lambda_context.tenant_id, + "durable": durable_seen, + "plain": plain_seen, + } + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("tenant-echo", tenant_echo) + runner.register_function("plain-echo", plain_tenant_echo) + result = runner.run(input=json.dumps({})) + with_tenant = runner.run(input=json.dumps({}), tenant_id="tenant-p") + + assert json.loads(result.result or "{}") == { + "parent": None, + "durable": None, + "plain": None, + } + # A tenant on run() reaches the parent; children invoked without one + # have none, as their invokes carried none. + assert json.loads(with_tenant.result or "{}") == { + "parent": "tenant-p", + "durable": None, + "plain": None, + } + + +def test_invoke_account_of_the_run_is_reported_to_handlers() -> None: + @durable_execution + def arn_echo(event: Any, context: DurableContext) -> str: + assert context.lambda_context is not None + return str(context.lambda_context.invoked_function_arn) + + def plain_arn_echo(event: Any, context: Any) -> str: + return str(context.invoked_function_arn) + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + durable_seen: Any = context.invoke("arn-echo", {}, name="durable") + plain_seen: Any = context.invoke("plain-arn-echo", {}, name="plain") + return {"durable": durable_seen, "plain": plain_seen} + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("arn-echo", arn_echo) + runner.register_function("plain-arn-echo", plain_arn_echo) + result = runner.run(input=json.dumps({}), account_id="999999999999") + + assert json.loads(result.result or "{}") == { + "durable": "arn:aws:lambda:us-west-2:999999999999:function:arn-echo", + "plain": "arn:aws:lambda:us-west-2:999999999999:function:plain-arn-echo", + } + + +def test_invoke_qualified_target_runs_the_qualified_registration_if_any() -> None: + """``child:prod`` runs a handler registered as ``child:prod``; with only + ``child`` registered, that single handler serves every qualifier.""" + + @durable_execution + def latest(event: Any, context: DurableContext) -> str: + return "latest" + + @durable_execution + def prod(event: Any, context: DurableContext) -> str: + return "prod" + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + return { + "prod": context.invoke("child:prod", {}, name="prod"), + "staging": context.invoke("child:staging", {}, name="staging"), + "bare": context.invoke("child", {}, name="bare"), + } + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_durable_function("child", latest) + runner.register_durable_function("child:prod", prod) + result = runner.run(input=json.dumps({})) + + assert json.loads(result.result or "{}") == { + "prod": "prod", + "staging": "latest", + "bare": "latest", + } + + +def test_plain_target_context_carries_the_requested_qualifier() -> None: + """A plain target invoked as ``child:prod`` sees a qualified function + ARN, whether it is registered under the qualified key or the bare name.""" + + def target(event: Any, context: Any) -> str: + return context.invoked_function_arn + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + return { + "exact": context.invoke("exact:prod", {}, name="exact"), + "fallback": context.invoke("bare:prod", {}, name="fallback"), + "bare": context.invoke("bare", {}, name="bare"), + "latest": context.invoke("bare:$LATEST", {}, name="latest"), + } + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_function("exact:prod", target) + runner.register_function("bare", target) + result = runner.run(input=json.dumps({})) + + prefix = "arn:aws:lambda:us-west-2:123456789012:function:" + assert json.loads(result.result or "{}") == { + "exact": prefix + "exact:prod", + "fallback": prefix + "bare:prod", + "bare": prefix + "bare", + "latest": prefix + "bare:$LATEST", + } + + +def test_targets_see_the_function_name_and_version_lambda_would_give() -> None: + """Lambda fills function_name, function_version and invoked_function_arn + on every invocation. Plain and durable targets, bare and qualified, + see the same three values here; an alias reports $LATEST because the + runner keeps no versions.""" + + def identity(context: Any) -> list: + return [ + context.function_name, + context.function_version, + context.invoked_function_arn.rsplit(":function:", 1)[1], + ] + + def plain(event: Any, context: Any) -> list: + return identity(context) + + @durable_execution + def durable(event: Any, context: DurableContext) -> list: + return identity(context.lambda_context) + + @durable_execution + def parent(event: Any, context: DurableContext) -> dict: + return { + "plain": context.invoke("plain", {}, name="p"), + "plain_version": context.invoke("plain:3", {}, name="pv"), + "durable": context.invoke("durable", {}, name="d"), + "durable_alias": context.invoke("durable:prod", {}, name="da"), + } + + with DurableFunctionTestRunner(handler=parent) as runner: + runner.register_function("plain", plain) + runner.register_durable_function("durable", durable) + result = runner.run(input=json.dumps({})) + + assert json.loads(result.result or "{}") == { + "plain": ["plain", "$LATEST", "plain"], + "plain_version": ["plain", "3", "plain:3"], + "durable": ["durable", "$LATEST", "durable"], + "durable_alias": ["durable", "$LATEST", "durable:prod"], + } diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py index 6418304d6..cd2fd8743 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/event_factory_test.py @@ -3,11 +3,13 @@ This module tests all the event creation factory methods in the Event class. """ +from dataclasses import replace from datetime import UTC, datetime from unittest.mock import Mock import pytest from aws_durable_execution_sdk_python.lambda_service import ( + ChainedInvokeOptions, ErrorObject, OperationStatus, OperationType, @@ -1873,21 +1875,23 @@ def test_from_operation_finished_no_details(self): # endregion from_operation_finished_tests -def test_chained_invoke_pending_details_from_dict(): - """Test ChainedInvokePendingDetails parsing in Event.from_dict.""" +def test_chained_invoke_started_details_full_from_dict(): + """Test ChainedInvokeStartedDetails parsing in Event.from_dict.""" data = { "EventType": "ChainedInvokeStarted", "EventTimestamp": datetime.now(UTC), - "ChainedInvokePendingDetails": { - "Input": {"Payload": "test-input", "Truncated": False}, + "ChainedInvokeStartedDetails": { "FunctionName": "test-function", + "Input": {"Payload": "test-input", "Truncated": False}, + "DurableExecutionArn": "child-arn", }, } event = Event.from_dict(data) - assert event.chained_invoke_pending_details is not None - assert event.chained_invoke_pending_details.input.payload == "test-input" - assert event.chained_invoke_pending_details.function_name == "test-function" + assert event.chained_invoke_started_details is not None + assert event.chained_invoke_started_details.function_name == "test-function" + assert event.chained_invoke_started_details.input.payload == "test-input" + assert event.chained_invoke_started_details.durable_execution_arn == "child-arn" def test_event_creation_context_sub_type_property(): @@ -2003,8 +2007,8 @@ def test_event_creation_context_get_retry_details(): assert retry_details is None -def test_create_chained_invoke_event_pending(): - """Test Event.create_chained_invoke_event_pending method.""" +def test_create_chained_invoke_event_started_from_start_update(): + """Test Event.create_chained_invoke_event_started with a START update.""" operation = Mock() operation.operation_id = "invoke-1" operation.name = "test_invoke" @@ -2013,6 +2017,16 @@ def test_create_chained_invoke_event_pending(): operation.start_timestamp = datetime.now(UTC) operation.sub_type = None + update = OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=OperationAction.START, + payload='{"n": 1}', + chained_invoke_options=ChainedInvokeOptions( + function_name="child-function", tenant_id=None + ), + ) + context = EventCreationContext.create( operation=operation, event_id=1, @@ -2025,13 +2039,61 @@ def test_create_chained_invoke_event_pending(): execution_timeout_seconds=300, execution_retention_period_days=7, ), + operation_update=update, include_execution_data=True, ) + context = replace(context, child_execution_arn="child-arn") - event = Event.create_chained_invoke_event_pending(context) + event = Event.create_chained_invoke_event_started(context) assert event.event_type == "ChainedInvokeStarted" assert event.operation_id == "invoke-1" assert event.name == "test_invoke" - assert event.chained_invoke_pending_details is not None - assert event.chained_invoke_pending_details.function_name == "test" + assert event.chained_invoke_started_details is not None + assert event.chained_invoke_started_details.function_name == "child-function" + assert event.chained_invoke_started_details.input.payload == '{"n": 1}' + assert event.chained_invoke_started_details.durable_execution_arn == "child-arn" + + +def test_create_chained_invoke_event_started_omits_the_tenant_id(): + """The service's history does not return TenantId for this event, so + the runner omits it too: a history assertion written locally must hold + against the service.""" + operation = Mock() + operation.operation_id = "invoke-1" + operation.name = "test_invoke" + operation.parent_id = None + operation.status = OperationStatus.PENDING + operation.start_timestamp = datetime.now(UTC) + operation.sub_type = None + + update = OperationUpdate( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + action=OperationAction.START, + payload="{}", + chained_invoke_options=ChainedInvokeOptions( + function_name="child-function", tenant_id="tenant-a" + ), + ) + context = EventCreationContext.create( + operation=operation, + event_id=1, + durable_execution_arn="arn:test", + start_input=StartDurableExecutionInput( + account_id="123", + function_name="test", + function_qualifier="$LATEST", + execution_name="test", + execution_timeout_seconds=300, + execution_retention_period_days=7, + ), + operation_update=update, + include_execution_data=True, + ) + + event = Event.create_chained_invoke_event_started(context) + + assert event.chained_invoke_started_details is not None + assert event.chained_invoke_started_details.tenant_id is None + assert "TenantId" not in event.to_dict()["ChainedInvokeStartedDetails"] diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py index 7ea66fe3c..bde5e3d4b 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/execution_test.py @@ -1,5 +1,7 @@ """Unit tests for execution module.""" +import json +from dataclasses import replace from datetime import datetime, timezone from unittest.mock import patch, Mock @@ -1220,6 +1222,9 @@ def test_execution_round_trip_preserves_new_fields(): execution.operation_last_touched_seq = {"op-A": 2, "op-B": 7} execution.operation_size_bytes = {"op-A": 42, "op-B": 100} execution.needs_reinvoke = True + execution.chained_invoke_children = {"invoke-1": "child-arn"} + execution.parent_execution_arn = "parent-arn" + execution.region = "eu-west-1" execution.last_checkpoint = CheckpointIdempotencyRecord( client_token="c1", inbound_checkpoint_token="tok-in", @@ -1230,6 +1235,9 @@ def test_execution_round_trip_preserves_new_fields(): rehydrated = Execution.from_json_dict(execution.to_json_dict()) + assert rehydrated.chained_invoke_children == {"invoke-1": "child-arn"} + assert rehydrated.parent_execution_arn == "parent-arn" + assert rehydrated.region == "eu-west-1" assert rehydrated.seq_counter == 7 assert rehydrated.token_sequence == 3 assert rehydrated.handler_seen_seq == 5 @@ -1270,6 +1278,9 @@ def test_execution_from_old_format_dict_uses_safe_defaults(): assert execution.operation_size_bytes == {} assert execution.needs_reinvoke is False assert execution.last_checkpoint is None + assert execution.chained_invoke_children == {} + assert execution.parent_execution_arn is None + assert execution.region is None def test_checkpoint_idempotency_record_equality(): @@ -1327,6 +1338,56 @@ def test_advance_token_sequence_returns_new_value_and_leaves_seq_counter(): # endregion # region OperationPaginatorState +def test_complete_chained_invoke_sizes_the_operation_by_its_result(): + execution = Execution( + "test-arn", + _make_start_input(), + [ + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=OperationStatus.STARTED, + ) + ], + ) + execution.operation_size_bytes["invoke-1"] = 3 + + execution.complete_chained_invoke( + "invoke-1", OperationStatus.SUCCEEDED, result="x" * 1000 + ) + assert execution.operation_size_bytes["invoke-1"] == 1000 + + execution.operations[0] = replace( + execution.operations[0], status=OperationStatus.STARTED + ) + error = ErrorObject.from_message("boom") + execution.complete_chained_invoke("invoke-1", OperationStatus.FAILED, error=error) + assert execution.operation_size_bytes["invoke-1"] == len( + json.dumps(error.to_dict()) + ) + + +def _chained_invoke_execution(status: OperationStatus) -> Execution: + return Execution( + "test-arn", + _make_start_input(), + [ + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=status, + ) + ], + ) + + +def test_complete_chained_invoke_rejects_a_terminal_operation(): + execution = _chained_invoke_execution(OperationStatus.SUCCEEDED) + + with pytest.raises(IllegalStateException, match="not active"): + execution.complete_chained_invoke("invoke-1", OperationStatus.FAILED) + + def _make_execution_with_ops( op_ids: list[str], sizes: dict[str, int] | None = None, @@ -1605,8 +1666,10 @@ def test_complete_wait_records_updated_operation_id(): assert execution.updated_operation_ids == ["wait-1"] -def test_record_invocation_completion_clears_updated_operation_ids(): - """Updated IDs are scoped to the next completed invocation.""" +def test_record_invocation_completion_keeps_updated_operation_ids(): + """An operation that changes while the handler runs is reported on the + next invocation. So completing the invocation keeps the list; only + delivering the state in an input resets it.""" start_input = StartDurableExecutionInput( account_id="123456789012", function_name="test-function", @@ -1621,5 +1684,26 @@ def test_record_invocation_completion_clears_updated_operation_ids(): now = datetime(2023, 1, 1, 12, 0, 0, tzinfo=timezone.utc) execution.record_invocation_completion(now, now, "request-1") + assert execution.updated_operation_ids == ["wait-1"] + execution.mark_state_delivered() assert execution.updated_operation_ids == [] + + +def test_function_arn_is_qualified_with_the_executed_version(): + execution = Execution.new(_make_start_input()) + execution.region = "ap-southeast-2" + + assert execution.executed_version() == "$LATEST" + assert ( + execution.function_arn("us-west-2") + == "arn:aws:lambda:ap-southeast-2:123456789012:function:test-function:$LATEST" + ) + + +def test_function_arn_falls_back_to_the_default_region_when_none_was_recorded(): + """An execution stored before the runner recorded regions has none.""" + execution = Execution.new(_make_start_input()) + assert execution.region is None + + assert execution.function_arn("us-west-2").startswith("arn:aws:lambda:us-west-2:") diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/executor_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/executor_test.py index b72921242..bb6b71f4f 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/executor_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/executor_test.py @@ -1,6 +1,7 @@ """Unit tests for executor module.""" import asyncio +import time from datetime import UTC, datetime from unittest.mock import ANY, Mock, patch @@ -28,11 +29,15 @@ InvalidParameterValueException, ResourceNotFoundException, ) +from aws_durable_execution_sdk_python_testing.child_dispatcher import KnownOutcome from aws_durable_execution_sdk_python_testing.execution import ( ExecutionStatus, Execution, ) -from aws_durable_execution_sdk_python_testing.executor import Executor, InvocationState +from aws_durable_execution_sdk_python_testing.executor import ( + Executor, + InvocationState, +) from aws_durable_execution_sdk_python_testing.invoker import InvokeResponse from aws_durable_execution_sdk_python_testing.model import ( ListDurableExecutionsResponse, @@ -45,6 +50,9 @@ from aws_durable_execution_sdk_python_testing.observer import ( ExecutionObserver, ) +from aws_durable_execution_sdk_python_testing.stores.filesystem import ( + FileSystemExecutionStore, +) from aws_durable_execution_sdk_python_testing.stores.memory import ( InMemoryExecutionStore, ) @@ -52,6 +60,9 @@ CallbackToken, CheckpointToken, ) +from aws_durable_execution_sdk_python_testing.worker.checkpoint_tasks import ( + CallableTask, +) class MockExecutionObserver(ExecutionObserver): @@ -63,6 +74,7 @@ def __init__(self): self.wait_timers = {} self.retry_schedules = {} self.callback_creations = {} + self.chained_invoke_starts = {} def on_completed(self, execution_arn: str, result: str | None = None) -> None: """Capture completion events.""" @@ -85,6 +97,22 @@ def on_callback_created( "callback_id": callback_token.to_str(), } + def on_chained_invoke_started( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> None: + """Capture chained invoke dispatch events.""" + self.chained_invoke_starts[execution_arn] = { + "operation_id": operation_id, + "function_name": function_name, + "tenant_id": tenant_id, + "payload": payload, + } + def on_callback_completed( self, execution_arn: str, operation_id: str, callback_id: str ) -> None: @@ -621,6 +649,8 @@ def test_should_retry_when_pending_response_has_no_operations( mock_execution.start_input = start_input mock_execution.consecutive_failed_invocation_attempts = 0 mock_execution.has_pending_operations.return_value = False # No pending operations + mock_execution.handler_seen_seq = 0 + mock_execution.has_changes_after.return_value = False # Mock invoker to return pending response mock_invocation_input = Mock() @@ -654,6 +684,128 @@ def test_should_retry_when_pending_response_has_no_operations( assert mock_scheduler.call_later.call_count == 3 +def test_pending_response_is_valid_when_an_operation_completed_after_the_handler_saw_it( + executor, mock_store, mock_scheduler, mock_invoker, start_input +): + """An operation that completes between the handler's return and the + response check leaves no pending operation, but a change after the + invocation's input was built: PENDING is accepted and the execution + is re-invoked instead of counted as a failed attempt.""" + mock_execution = Mock() + mock_execution.durable_execution_arn = "test-arn" + mock_execution.is_complete = False + mock_execution.start_input = start_input + mock_execution.consecutive_failed_invocation_attempts = 0 + mock_execution.seq_counter = 4 + mock_execution.handler_seen_seq = 2 + mock_execution.has_pending_operations.return_value = False + mock_execution.has_changes_after.side_effect = lambda seq: seq == 4 + mock_execution.has_unseen_changes.return_value = True + mock_execution.needs_reinvoke = True + + mock_invoker.create_invocation_input.return_value = Mock() + mock_invoker.invoke.return_value = InvokeResponse( + invocation_output=DurableExecutionInvocationOutput( + status=InvocationStatus.PENDING + ), + request_id="test-request-id", + ) + + with patch( + "aws_durable_execution_sdk_python_testing.executor.Execution" + ) as mock_execution_class: + mock_execution_class.new.return_value = mock_execution + mock_store.load.return_value = mock_execution + + executor.start_execution(start_input) + handler = mock_scheduler.call_later.call_args_list[-1][0][0] + asyncio.run(handler()) + + assert mock_execution.consecutive_failed_invocation_attempts == 0 + # timeout + initial invocation + the re-invoke, with no retry delay + assert mock_scheduler.call_later.call_count == 3 + assert mock_scheduler.call_later.call_args_list[-1].kwargs["delay"] == 0 + + +def test_begin_invocation_baseline_predates_a_completion_queued_behind_it( + executor, mock_store, mock_invoker, start_input +): + """The baseline is the counter when the input was built. A completion + applied to the same object afterward lies past it.""" + execution = Execution.new(start_input) + execution.operations.append( + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=OperationStatus.STARTED, + ) + ) + mock_store.load.return_value = execution + mock_invoker.create_invocation_input.return_value = Mock() + + claim = executor._begin_invocation(execution.durable_execution_arn) # noqa: SLF001 + assert claim is not None + _, _, baseline = claim + assert baseline == execution.seq_counter + + execution.complete_chained_invoke("invoke-1", OperationStatus.SUCCEEDED, result="1") + assert execution.seq_counter > baseline + assert execution.has_changes_after(baseline) + + +def test_pending_response_with_only_input_time_state_is_still_an_error( + executor, mock_store, start_input +): + """Operations the invocation input already delivered lie past the + handler's watermark but not past the invocation's baseline, so a + no-pending PENDING is rejected rather than accepted without a re-invoke.""" + execution = Execution.new(start_input) + execution.operations.append( + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=OperationStatus.STARTED, + ) + ) + execution.complete_chained_invoke("invoke-1", OperationStatus.SUCCEEDED, result="1") + assert execution.has_unseen_changes() + baseline = execution.seq_counter + + with pytest.raises(InvalidParameterValueException, match="no pending operations"): + executor._validate_invocation_response_and_store( # noqa: SLF001 + execution.durable_execution_arn, + DurableExecutionInvocationOutput(status=InvocationStatus.PENDING), + execution, + baseline, + ) + + +def test_pending_response_is_an_error_when_the_handler_saw_the_completion( + executor, mock_store, start_input +): + """A completion after the baseline that a checkpoint response already + delivered leaves nothing to re-invoke for, so PENDING is rejected.""" + execution = Execution.new(start_input) + execution.operations.append( + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=OperationStatus.STARTED, + ) + ) + baseline = execution.seq_counter + execution.complete_chained_invoke("invoke-1", OperationStatus.SUCCEEDED, result="1") + execution.handler_seen_seq = execution.seq_counter + + with pytest.raises(InvalidParameterValueException, match="no pending operations"): + executor._validate_invocation_response_and_store( # noqa: SLF001 + execution.durable_execution_arn, + DurableExecutionInvocationOutput(status=InvocationStatus.PENDING), + execution, + baseline, + ) + + def test_invoke_handler_success( executor, mock_store, mock_scheduler, mock_invoker, start_input ): @@ -663,6 +815,7 @@ def test_invoke_handler_success( mock_execution.durable_execution_arn = "test-arn" mock_execution.is_complete = False mock_execution.start_input = start_input + mock_execution.region = "us-west-2" mock_invocation_input = Mock() mock_invoker.create_invocation_input.return_value = mock_invocation_input @@ -695,7 +848,12 @@ def test_invoke_handler_success( execution=mock_execution ) mock_invoker.invoke.assert_called_once_with( - "test-function", mock_invocation_input, None + "test-function", + mock_invocation_input, + None, + tenant_id=None, + account_id="123456789012", + region_name="us-west-2", ) @@ -815,12 +973,13 @@ def test_invoke_handler_resource_not_found( # Assert - verify workflow failure was triggered through public API mock_fail.assert_called_once() - # Verify the error contains the expected message + # Verify the error keeps the Lambda API error code and message, so a + # parent chained invoke sees ResourceNotFoundException. call_args = mock_fail.call_args assert call_args[0][0] == "test-arn" # execution_arn is first positional arg - assert "Function not found" in str( - call_args[0][1] - ) # error is second positional arg + error = call_args[0][1] + assert error.type == "ResourceNotFoundException" + assert error.message == "Function not found" def test_invoke_handler_general_exception( @@ -1189,6 +1348,7 @@ def test_complete_events_through_complete_execution( mock_execution = Mock() mock_execution.result = "test result" mock_store.load.return_value = mock_execution + mock_execution.parent_execution_arn = None # top-level: no parent to notify # Set up completion event through start_execution mock_event = Mock() @@ -1222,6 +1382,7 @@ def test_complete_events_no_event_through_public_api(executor, mock_store): mock_execution = Mock() mock_execution.result = "test result" mock_store.load.return_value = mock_execution + mock_execution.parent_execution_arn = None # top-level: no parent to notify # Complete execution without setting up completion event first # Should not raise exception when event doesn't exist @@ -1415,6 +1576,7 @@ def test_invoke_handler_execution_completed_during_invocation_async( incomplete_execution.start_input = start_input incomplete_execution.consecutive_failed_invocation_attempts = 0 incomplete_execution.durable_execution_arn = "test-arn" + incomplete_execution.seq_counter = 0 completed_execution = Mock(spec=Execution) completed_execution.is_complete = True @@ -2471,6 +2633,7 @@ def test_send_callback_heartbeat_invalid_token(executor): def test_complete_events_no_event(executor): """Test _complete_events when no event exists.""" + executor._store.load.return_value.parent_execution_arn = None # top-level # Should not raise exception when event doesn't exist executor._complete_events("nonexistent-arn") # Should handle gracefully @@ -2767,3 +2930,644 @@ def test_on_stopped(executor): executor.on_stopped("test-arn", error) mock_fail.assert_called_once_with("test-arn", error) + + +def test_resolve_dispatch_rejects_a_malformed_target_it_was_handed( + executor, mock_store, start_input +): + """Checkpoint validation normally rejects a malformed target first; a + dispatch that bypassed it still fails with the service's message.""" + parent = Mock() + parent.is_complete = False + parent.start_input = start_input + mock_store.load.return_value = parent + executor._child_dispatcher = Mock() # noqa: SLF001 + + result = executor._resolve_dispatch( # noqa: SLF001 + "parent-arn", "op-1", "not a function", None, None + ) + + assert isinstance(result, KnownOutcome) + assert result.outcome.error is not None + assert result.outcome.error.message == "Invalid function ARN 'not a function'" + executor._child_dispatcher.dispatch.assert_not_called() # noqa: SLF001 + + +def test_resolve_dispatch_carries_name_and_qualifier_separately( + executor, mock_store, start_input +): + parent = Mock() + parent.is_complete = False + parent.start_input = start_input + mock_store.load.return_value = parent + executor._child_dispatcher = Mock() # noqa: SLF001 + + executor._resolve_dispatch( # noqa: SLF001 + "parent-arn", + "op-1", + "arn:aws:lambda:us-west-2:123456789012:function:child:prod", + None, + "{}", + ) + + request = executor._child_dispatcher.dispatch.call_args.args[0] # noqa: SLF001 + assert (request.function_name, request.qualifier) == ("child", "prod") + assert request.lookup_keys() == ("child:prod", "child") + + +# region Chained invoke preflight + + +def test_preflight_without_a_dispatcher_fails_the_target(executor): + """A runner with no dispatcher cannot dispatch anything, so the target + fails in the checkpoint response with the same message dispatch gives.""" + executor._child_dispatcher = None # noqa: SLF001 + + error = executor.preflight_chained_invoke("child") + + assert error is not None + assert error.message == "This runner has no chained-invoke dispatcher configured." + assert error.type is None + + +def test_preflight_rejects_a_malformed_target(executor): + executor._child_dispatcher = Mock() # noqa: SLF001 + + error = executor.preflight_chained_invoke("not a function name!") + + assert error is not None + assert "not a function name!" in (error.message or "") + executor._child_dispatcher.preflight.assert_not_called() # noqa: SLF001 + + +def test_preflight_asks_the_dispatcher_with_the_parsed_target(executor): + from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + ChildOutcome, + FunctionTarget, + ) + + dispatcher = Mock() + dispatcher.preflight.return_value = None + executor._child_dispatcher = dispatcher # noqa: SLF001 + + assert executor.preflight_chained_invoke("child:prod") is None + dispatcher.preflight.assert_called_once_with( + FunctionTarget(name="child", qualifier="prod", account_id=None, region=None) + ) + + error = ErrorObject.from_message("Function not found: child.") + dispatcher.preflight.return_value = ChildOutcome(error=error) + assert executor.preflight_chained_invoke("child") is error + + +# endregion + + +def _linked_parent_and_child( + store, + parent_name: str = "parent", + child_name: str = "child", + *, + child_terminal: bool = True, +) -> tuple[Execution, Execution]: + """Persist a parent with a STARTED chained invoke and its linked child. + + Both halves of the link are stored: the parent lists the child under + the operation id, and the child names its parent. The child has + succeeded unless ``child_terminal`` is False, in which case it is + still running. + """ + parent = Execution.new( + StartDurableExecutionInput( + account_id="123456789012", + function_name=parent_name, + function_qualifier="$LATEST", + execution_name="parent-exec", + execution_timeout_seconds=300, + execution_retention_period_days=7, + invocation_id="parent-inv", + ) + ) + parent.start() + parent.operations.append( + Operation( + operation_id="invoke-1", + operation_type=OperationType.CHAINED_INVOKE, + status=OperationStatus.STARTED, + ) + ) + child = Execution.new( + StartDurableExecutionInput( + account_id="123456789012", + function_name=child_name, + function_qualifier="$LATEST", + execution_name="child-exec", + execution_timeout_seconds=300, + execution_retention_period_days=7, + invocation_id="child-inv", + ) + ) + child.parent_execution_arn = parent.durable_execution_arn + parent.record_chained_invoke_child("invoke-1", child.durable_execution_arn) + child.start() + if child_terminal: + child.complete_success('{"squared": 16}') + store.save(parent) + store.save(child) + return parent, child + + +def test_terminal_child_reaches_its_parent_after_a_restart( + tmp_path, mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """A fresh executor over the same store has no memory link for the + child. The child is then stopped through the public API, as a + StopDurableExecution call would after a restart. The persisted halves + of the link are enough: the parent's operation completes with the + stop error and the parent is re-invoked. + """ + store = FileSystemExecutionStore(tmp_path) + parent, child = _linked_parent_and_child(store, child_terminal=False) + restarted = Executor(store, mock_scheduler, mock_invoker, mock_checkpoint_processor) + assert restarted._chained_invoke_links == {} # noqa: SLF001 + + restarted.stop_execution( + child.durable_execution_arn, + error=ErrorObject.from_message("operator stopped the child"), + ) + # The parent-side transition runs on the parent's lane, in FIFO + # order; a no-op behind it is a barrier. + restarted._registry.submit( # noqa: SLF001 + parent.durable_execution_arn, CallableTask(lambda: None) + ).result(timeout=5) + + _, operation = store.load(parent.durable_execution_arn).find_operation("invoke-1") + assert operation.status is OperationStatus.STOPPED + assert ( + operation.chained_invoke_details.error.message == "operator stopped the child" + ) + assert ( + store.load(child.durable_execution_arn).current_status() + is ExecutionStatus.STOPPED + ) + # The parent is re-invoked to observe the completion. + assert mock_scheduler.call_later.call_count >= 1 + + +def test_persisted_link_completes_the_parent_when_the_child_already_succeeded( + tmp_path, mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """The fallback itself, with a child that succeeded: the parent's + operation takes the child's result.""" + store = FileSystemExecutionStore(tmp_path) + parent, child = _linked_parent_and_child(store) + restarted = Executor(store, mock_scheduler, mock_invoker, mock_checkpoint_processor) + + restarted._notify_parent_of_terminal_child(child.durable_execution_arn) # noqa: SLF001 + restarted._registry.submit( # noqa: SLF001 + parent.durable_execution_arn, CallableTask(lambda: None) + ).result(timeout=5) + + _, operation = store.load(parent.durable_execution_arn).find_operation("invoke-1") + assert operation.status is OperationStatus.SUCCEEDED + assert operation.chained_invoke_details.result == '{"squared": 16}' + + +def test_terminal_execution_without_a_parent_notifies_nobody( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + store = InMemoryExecutionStore() + parent, child = _linked_parent_and_child(store) + child.parent_execution_arn = None + store.save(child) + fresh = Executor(store, mock_scheduler, mock_invoker, mock_checkpoint_processor) + + fresh._notify_parent_of_terminal_child(child.durable_execution_arn) # noqa: SLF001 + + _, operation = store.load(parent.durable_execution_arn).find_operation("invoke-1") + assert operation.status is OperationStatus.STARTED + mock_scheduler.call_later.assert_not_called() + + +def test_persisted_link_ignores_a_parent_that_does_not_list_the_child( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + store = InMemoryExecutionStore() + parent, child = _linked_parent_and_child(store) + parent.chained_invoke_children.clear() + store.save(parent) + fresh = Executor(store, mock_scheduler, mock_invoker, mock_checkpoint_processor) + + assert fresh._persisted_link(child) is None # noqa: SLF001 + assert fresh._persisted_link(None) is None # noqa: SLF001 + + +@pytest.mark.parametrize( + ("qualifier", "arn_suffix", "version"), + [ + ("$LATEST", "child:$LATEST", "$LATEST"), + ("3", "child:3", "3"), + ("prod", "child:$LATEST", "$LATEST"), + ("$LATEST.PUBLISHED", "child:$LATEST.PUBLISHED", "$LATEST.PUBLISHED"), + ], +) +def test_get_execution_details_reports_the_qualified_identity( + mock_scheduler, + mock_invoker, + mock_checkpoint_processor, + qualifier, + arn_suffix, + version, +): + """FunctionArn is qualified with the executed version and Version + carries the same value, as the service reports them; both come from + the execution's own region.""" + store = InMemoryExecutionStore() + real_executor = Executor( + store, + mock_scheduler, + mock_invoker, + mock_checkpoint_processor, + region="us-west-2", + ) + output = real_executor.start_execution( + StartDurableExecutionInput( + account_id="123456789012", + function_name="child", + function_qualifier=qualifier, + execution_name="child-run", + execution_timeout_seconds=300, + execution_retention_period_days=7, + invocation_id="inv-1", + ) + ) + + details = real_executor.get_execution_details(output.execution_arn) + + assert ( + details.function_arn + == f"arn:aws:lambda:us-west-2:123456789012:function:{arn_suffix}" + ) + assert details.version == version + assert real_executor.region == "us-west-2" + + +def _run_dispatch_with_blocked_target( + executor, *, shutdown_before_result: bool, block_in: str = "invocation" +): + """Drive one dispatch that blocks until released, in target resolution + or in the plain-target invoke. + + Returns the number of completions applied to the parent. + """ + import threading + + from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + ChildOutcome, + RunInvocation, + ) + + release = threading.Event() + applied: list[ChildOutcome] = [] + + def invocation() -> ChildOutcome: + if block_in == "invocation": + release.wait(5) + return ChildOutcome(result="42") + + def resolve(*_args) -> RunInvocation: + if block_in == "resolve": + release.wait(5) + return RunInvocation(invocation=invocation) + + async def record(_arn: str, _op: str, outcome: ChildOutcome) -> None: + applied.append(outcome) + + executor._resolve_dispatch = resolve # noqa: SLF001 + executor._finish_from_outcome = record # noqa: SLF001 + dispatch = executor._dispatch_chained_invoke( # noqa: SLF001 + "parent-arn", "invoke-1", "child", None, None + ) + + async def main() -> None: + task = asyncio.ensure_future(dispatch()) + await asyncio.sleep(0.1) # the blocking thread is now waiting + if shutdown_before_result: + executor.shutdown() + release.set() + await task + + asyncio.run(main()) + return len(applied) + + +def test_a_target_result_that_lands_after_shutdown_is_dropped(executor): + """The blocked target thread returns after shutdown. Its result must + not be applied to the parent, whose store may already be closed.""" + assert _run_dispatch_with_blocked_target(executor, shutdown_before_result=True) == 0 + + +def test_a_target_resolved_after_shutdown_is_not_dispatched(executor): + """Shutdown while the target is still being resolved: nothing is + invoked and nothing is applied.""" + assert ( + _run_dispatch_with_blocked_target( + executor, shutdown_before_result=True, block_in="resolve" + ) + == 0 + ) + + +def test_a_target_result_before_shutdown_is_applied(executor): + """Positive control for the tests above: without shutdown the same + result completes the parent's operation once.""" + assert ( + _run_dispatch_with_blocked_target(executor, shutdown_before_result=False) == 1 + ) + + +def test_a_handler_response_that_lands_after_shutdown_is_not_recorded( + executor, mock_store, mock_scheduler, mock_invoker, start_input +): + """The handler invocation returns after shutdown. Its response is + dropped instead of being recorded on the execution.""" + import threading + + mock_execution = Mock() + mock_execution.durable_execution_arn = "test-arn" + mock_execution.is_complete = False + mock_execution.start_input = start_input + mock_execution.region = "us-west-2" + mock_invoker.create_invocation_input.return_value = Mock() + release = threading.Event() + + def blocked_invoke(*_args, **_kwargs) -> InvokeResponse: + release.wait(5) + return InvokeResponse( + invocation_output=DurableExecutionInvocationOutput( + status=InvocationStatus.SUCCEEDED, result="late" + ), + request_id="late-request", + ) + + mock_invoker.invoke.side_effect = blocked_invoke + executor._finish_invocation = Mock() # noqa: SLF001 + + with patch( + "aws_durable_execution_sdk_python_testing.executor.Execution" + ) as mock_execution_class: + mock_execution_class.new.return_value = mock_execution + mock_store.load.return_value = mock_execution + executor.start_execution(start_input) + handler = mock_scheduler.call_later.call_args_list[-1][0][0] + + async def main() -> None: + task = asyncio.ensure_future(handler()) + await asyncio.sleep(0.1) # the invoke thread is now blocked + executor.shutdown() + release.set() + await task + + asyncio.run(main()) + + mock_invoker.invoke.assert_called_once() + executor._finish_invocation.assert_not_called() # noqa: SLF001 + + +def test_a_child_whose_record_cannot_be_read_fails_the_parent_operation( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """The memory link names the parent, but the child cannot be loaded. + The parent's operation still completes, as FAILED, so the parent is + never left waiting on a child the runner cannot read.""" + store = InMemoryExecutionStore() + parent, child = _linked_parent_and_child(store) + real_executor = Executor( + store, mock_scheduler, mock_invoker, mock_checkpoint_processor + ) + real_executor._chained_invoke_links[child.durable_execution_arn] = ( # noqa: SLF001 + parent.durable_execution_arn, + "invoke-1", + ) + store._store.pop(child.durable_execution_arn) # noqa: SLF001 + + real_executor._notify_parent_of_terminal_child(child.durable_execution_arn) # noqa: SLF001 + real_executor._registry.submit( # noqa: SLF001 + parent.durable_execution_arn, CallableTask(lambda: None) + ).result(timeout=5) + + _, operation = store.load(parent.durable_execution_arn).find_operation("invoke-1") + assert operation.status is OperationStatus.FAILED + assert "could not be read" in operation.chained_invoke_details.error.message + + +def test_a_child_start_that_runs_into_shutdown_persists_and_launches_nothing( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """Shutdown begins while the child-start task is already running. + Nothing is persisted and nothing is scheduled: a child created now + would never be invoked.""" + store = InMemoryExecutionStore() + parent, _ = _linked_parent_and_child(store) + real_executor = Executor( + store, mock_scheduler, mock_invoker, mock_checkpoint_processor + ) + child_start = StartDurableExecutionInput( + account_id="123456789012", + function_name="child", + function_qualifier="$LATEST", + execution_name="late-child", + execution_timeout_seconds=300, + execution_retention_period_days=7, + ) + real_executor.shutdown() + + real_executor._start_child_execution( # noqa: SLF001 + parent.durable_execution_arn, "invoke-2", child_start + ) + + assert len(store.list_all()) == 2 # parent and the original child only + mock_scheduler.call_later.assert_not_called() + + +def test_a_completion_during_the_invocation_reaches_the_next_input( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """A chained target completes while the parent handler is still + running. The service reports it in the next invocation's + UpdatedOperationIds, because the handler never observed it. So must + the runner: completing the current invocation keeps the id, and the + next input carries it.""" + store = InMemoryExecutionStore() + parent, _child = _linked_parent_and_child(store) + real_executor = Executor( + store, mock_scheduler, mock_invoker, mock_checkpoint_processor + ) + arn = parent.durable_execution_arn + now = datetime.now(UTC) + + # The parent handler is running: the child's completion lands first. + real_executor._apply_chained_invoke_completion( # noqa: SLF001 + arn, "invoke-1", OperationStatus.SUCCEEDED, '{"squared": 16}', None + ) + assert store.load(arn).updated_operation_ids == ["invoke-1"] + + # Then the current invocation returns. + store.load(arn).record_invocation_completion(now, now, "request-1") + store.save(store.load(arn)) + assert store.load(arn).updated_operation_ids == ["invoke-1"] + + # The next input carries the id, and delivering it resets the list. + execution = store.load(arn) + execution.begin_new_invocation() + invocation_input = mock_invoker.create_invocation_input(execution=execution) + mock_invoker.create_invocation_input.assert_called_once() + assert invocation_input is mock_invoker.create_invocation_input.return_value + execution.mark_state_delivered() + assert execution.updated_operation_ids == [] + + +def test_a_target_exception_that_lands_after_shutdown_is_dropped(executor): + """The target raises after shutdown, or the pool refuses it. The + exception path must not complete the parent's operation either.""" + import threading + + from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + ChildOutcome, + RunInvocation, + ) + + release = threading.Event() + applied: list[ChildOutcome] = [] + + def invocation() -> ChildOutcome: + release.wait(5) + msg = "target exploded after shutdown" + raise RuntimeError(msg) + + async def record(_arn: str, _op: str, outcome: ChildOutcome) -> None: + applied.append(outcome) + + executor._resolve_dispatch = Mock( # noqa: SLF001 + return_value=RunInvocation(invocation=invocation) + ) + executor._finish_from_outcome = record # noqa: SLF001 + dispatch = executor._dispatch_chained_invoke( # noqa: SLF001 + "parent-arn", "invoke-1", "child", None, None + ) + + async def main() -> None: + task = asyncio.ensure_future(dispatch()) + await asyncio.sleep(0.1) + executor.shutdown() + release.set() + await task + + asyncio.run(main()) + assert applied == [] + + +def test_a_handler_error_that_lands_after_shutdown_is_not_retried( + executor, mock_store, mock_scheduler, mock_invoker, start_input +): + """The handler invocation raises after shutdown. No retry is recorded + and nothing is rescheduled.""" + import threading + + mock_execution = Mock() + mock_execution.durable_execution_arn = "test-arn" + mock_execution.is_complete = False + mock_execution.start_input = start_input + mock_execution.region = "us-west-2" + mock_invoker.create_invocation_input.return_value = Mock() + release = threading.Event() + + def failing_invoke(*_args, **_kwargs) -> InvokeResponse: + release.wait(5) + msg = "connection reset" + raise ConnectionError(msg) + + mock_invoker.invoke.side_effect = failing_invoke + executor._retry_after_error = Mock() # noqa: SLF001 + + with patch( + "aws_durable_execution_sdk_python_testing.executor.Execution" + ) as mock_execution_class: + mock_execution_class.new.return_value = mock_execution + mock_store.load.return_value = mock_execution + executor.start_execution(start_input) + handler = mock_scheduler.call_later.call_args_list[-1][0][0] + + async def main() -> None: + task = asyncio.ensure_future(handler()) + await asyncio.sleep(0.1) + executor.shutdown() + release.set() + await task + + asyncio.run(main()) + + executor._retry_after_error.assert_not_called() # noqa: SLF001 + + +def test_shutdown_waits_for_a_child_start_already_past_its_check( + mock_scheduler, mock_invoker, mock_checkpoint_processor +): + """Shutdown begins while a child start is blocked inside + _create_execution, after its check. Shutdown waits for that start, + so the child is created, linked and launched as one unit, and no + half-made child is left in the store.""" + import threading + + store = InMemoryExecutionStore() + parent, _ = _linked_parent_and_child(store) + real_executor = Executor( + store, mock_scheduler, mock_invoker, mock_checkpoint_processor + ) + child_start = StartDurableExecutionInput( + account_id="123456789012", + function_name="child", + function_qualifier="$LATEST", + execution_name="racing-child", + execution_timeout_seconds=300, + execution_retention_period_days=7, + ) + inside_create = threading.Event() + release_create = threading.Event() + original_create = real_executor._create_execution # noqa: SLF001 + + def blocked_create(*args, **kwargs): + inside_create.set() + release_create.wait(5) + return original_create(*args, **kwargs) + + real_executor._create_execution = blocked_create # noqa: SLF001 + + starter = threading.Thread( + target=real_executor._start_child_execution, # noqa: SLF001 + args=(parent.durable_execution_arn, "invoke-2", child_start), + ) + starter.start() + assert inside_create.wait(5) + + shutdown_done = threading.Event() + threading.Thread( + target=lambda: (real_executor.shutdown(), shutdown_done.set()) + ).start() + time.sleep(0.2) + assert not shutdown_done.is_set() # shutdown is waiting for the start + + release_create.set() + starter.join(5) + assert shutdown_done.wait(5) + + children = [ + e for e in store.list_all() if e.start_input.execution_name == "racing-child" + ] + assert len(children) == 1 + assert ( + store.load(parent.durable_execution_arn).chained_invoke_children["invoke-2"] + == children[0].durable_execution_arn + ) + # The child was launched: its timeout and its first invocation are scheduled. + assert mock_scheduler.call_later.call_count == 2 diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/invoker_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/invoker_test.py index 02e7b3013..14e6f2344 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/invoker_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/invoker_test.py @@ -1,9 +1,11 @@ """Tests for invoker module.""" +import io import json from unittest.mock import Mock, patch import pytest +from botocore.awsrequest import AWSResponse # type: ignore from aws_durable_execution_sdk_python.execution import ( DurableExecutionInvocationInput, DurableExecutionInvocationInputWithClient, @@ -29,10 +31,15 @@ OperationPaginatorState, ) from aws_durable_execution_sdk_python_testing.invoker import ( + DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + INVOCATION_MARKER_HEADER, + INVOCATION_MARKER_VALUE, + LAMBDA_READ_TIMEOUT_HEADROOM_SECONDS, InProcessInvoker, LambdaInvoker, - _LAMBDA_CLIENT_CONFIG, + create_lambda_client, create_test_lambda_context, + read_timeout_for, ) from aws_durable_execution_sdk_python_testing.model import ( LambdaContext, @@ -48,10 +55,154 @@ def test_create_test_lambda_context(): context.invoked_function_arn == "arn:aws:lambda:us-west-2:123456789012:function:test-function" ) - assert context.tenant_id == "test-tenant-789" + assert context.tenant_id is None # no tenant given, as in Lambda assert context.client_context is not None +def test_create_test_lambda_context_reports_the_execution_values(): + context = create_test_lambda_context( + region="eu-west-1", + account_id="999999999999", + function_name="child", + tenant_id="tenant-a", + ) + assert ( + context.invoked_function_arn + == "arn:aws:lambda:eu-west-1:999999999999:function:child" + ) + assert context.tenant_id == "tenant-a" + + +@pytest.mark.parametrize( + ("identifier", "name", "version", "arn_suffix"), + [ + ("child", "child", "$LATEST", ":function:child"), + ("child:$LATEST", "child", "$LATEST", ":function:child:$LATEST"), + ("child:7", "child", "7", ":function:child:7"), + ("child:prod", "child", "$LATEST", ":function:child:prod"), + ( + "child:$LATEST.PUBLISHED", + "child", + "$LATEST.PUBLISHED", + ":function:child:$LATEST.PUBLISHED", + ), + ], +) +def test_create_test_lambda_context_fills_name_version_and_arn( + identifier, name, version, arn_suffix +): + """Lambda fills all three identity fields; target code reading them + must see the same values locally. An alias's version is not known to + the runner, so it reports $LATEST.""" + context = create_test_lambda_context(function_name=identifier) + + assert context.function_name == name + assert context.function_version == version + assert context.invoked_function_arn.endswith(arn_suffix) + + +def test_in_process_invoker_hands_the_function_identity_to_a_durable_handler(): + seen: dict = {} + + def handler(event, context): # noqa: ARG001 + seen.update( + name=context.function_name, + version=context.function_version, + arn=context.invoked_function_arn, + ) + return {"Status": "SUCCEEDED", "Result": "ok"} + + invoker = InProcessInvoker(handler, Mock(), region="us-west-2") + invoker.register("child:prod", handler) + + invoker.invoke("child:prod", _minimal_input()) + assert seen == { + "name": "child", + "version": "$LATEST", + "arn": "arn:aws:lambda:us-west-2:123456789012:function:child:prod", + } + + invoker.invoke("child", _minimal_input()) + assert (seen["name"], seen["version"]) == ("child", "$LATEST") + assert seen["arn"].endswith(":function:child") + + +def _minimal_input() -> DurableExecutionInvocationInput: + return DurableExecutionInvocationInput( + durable_execution_arn="test-arn", + checkpoint_token="test-token", # noqa: S106 + initial_execution_state=InitialExecutionState(operations=[], next_marker=""), + ) + + +def test_in_process_invoker_hands_the_tenant_and_region_to_the_handler(): + seen: dict = {} + + def handler(event, context): # noqa: ARG001 + seen["tenant_id"] = context.tenant_id + seen["arn"] = context.invoked_function_arn + return {"Status": "SUCCEEDED", "Result": "ok"} + + invoker = InProcessInvoker(handler, Mock(), region="eu-west-1") + invoker.invoke( + "child", _minimal_input(), tenant_id="tenant-a", account_id="999999999999" + ) + assert seen == { + "tenant_id": "tenant-a", + "arn": "arn:aws:lambda:eu-west-1:999999999999:function:child", + } + + invoker.invoke("child", _minimal_input()) + assert seen["tenant_id"] is None # no tenant: none, as in Lambda + assert seen["arn"] == "arn:aws:lambda:eu-west-1:123456789012:function:child" + + +def test_in_process_invoker_resolves_a_qualified_identifier_then_the_bare_name(): + seen: list = [] + + def make(tag): + def handler(event, context): # noqa: ARG001 + seen.append((tag, context.invoked_function_arn)) + return {"Status": "SUCCEEDED", "Result": tag} + + return handler + + invoker = InProcessInvoker(make("root"), Mock(), region="us-west-2") + invoker.register("child", make("latest")) + invoker.register("child:prod", make("prod")) + + invoker.invoke("child:prod", _minimal_input()) + invoker.invoke("child:staging", _minimal_input()) + invoker.invoke("child", _minimal_input()) + invoker.invoke("other", _minimal_input()) + assert [tag for tag, _ in seen] == ["prod", "latest", "latest", "root"] + assert seen[0][1].endswith(":function:child:prod") + + +def _stub_success_client() -> Mock: + client = Mock() + payload = Mock() + payload.read.return_value = json.dumps({"Status": "SUCCEEDED"}).encode("utf-8") + client.invoke.return_value = { + "StatusCode": 200, + "Payload": payload, + "ResponseMetadata": {"HTTPHeaders": {}}, + } + return client + + +def test_lambda_invoker_sends_tenant_id_when_the_execution_has_one(): + client = _stub_success_client() + LambdaInvoker(client).invoke("child", _minimal_input(), tenant_id="tenant-a") + assert client.invoke.call_args.kwargs["TenantId"] == "tenant-a" + + +def test_lambda_invoker_omits_tenant_id_when_the_execution_has_none(): + client = _stub_success_client() + LambdaInvoker(client).invoke("child", _minimal_input()) + assert "TenantId" not in client.invoke.call_args.kwargs + + def test_in_process_invoker_init(): """Test InProcessInvoker initialization.""" handler = Mock() @@ -139,12 +290,289 @@ def test_lambda_invoker_create(): assert isinstance(invoker, LambdaInvoker) assert invoker.lambda_client is mock_client - mock_boto3.client.assert_called_once_with( - "lambda", - endpoint_url="http://localhost:3001", - region_name="us-west-2", - config=_LAMBDA_CLIENT_CONFIG, + mock_boto3.client.assert_called_once() + kwargs = mock_boto3.client.call_args.kwargs + assert mock_boto3.client.call_args.args == ("lambda",) + assert kwargs["endpoint_url"] == "http://localhost:3001" + assert kwargs["region_name"] == "us-west-2" + assert kwargs["config"].read_timeout == DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS + assert kwargs["config"].retries == {"max_attempts": 0} + + +def test_read_timeout_outlasts_one_invocation(): + """The client waits one emulated function timeout plus headroom.""" + assert read_timeout_for(900) == 900 + LAMBDA_READ_TIMEOUT_HEADROOM_SECONDS + assert DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS == 960 + assert read_timeout_for(5400) == 5460 + + +def test_create_lambda_client_applies_read_timeout(): + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.return_value = Mock() + create_lambda_client("http://localhost:3001", "us-west-2", 5460) + assert mock_boto3.client.call_args.kwargs["config"].read_timeout == 5460 + + +def _clients_created(mock_boto3) -> list[tuple[str, str]]: + """(endpoint, region) of every client boto3 was asked for, in order.""" + return [ + (c.kwargs["endpoint_url"], c.kwargs["region_name"]) + for c in mock_boto3.client.call_args_list + ] + + +def test_lambda_invoker_same_endpoint_in_a_new_region_recreates_both_client_kinds(): + """The region is the signing region, so the same URL in another region + is another client, for handler invocations and for chained targets.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.side_effect = lambda *_a, **_k: Mock() + invoker = LambdaInvoker.create("http://localhost:3001", "us-west-2") + handler_before = invoker._get_client_for_execution("arn-1") # noqa: SLF001 + unmarked_before = invoker.unmarked_client_for("arn-1") + + invoker.update_endpoint("http://localhost:3001", "eu-west-1") + + handler_after = invoker._get_client_for_execution("arn-2") # noqa: SLF001 + unmarked_after = invoker.unmarked_client_for("arn-2") + assert handler_after is invoker.lambda_client + assert handler_after is not handler_before + assert unmarked_after is not unmarked_before + assert _clients_created(mock_boto3) == [ + ("http://localhost:3001", "us-west-2"), # handler client + ("http://localhost:3001", "us-west-2"), # unmarked client + ("http://localhost:3001", "eu-west-1"), # handler client + ("http://localhost:3001", "eu-west-1"), # unmarked client + ] + # One client per (endpoint, region) and kind. + assert invoker._get_client_for_execution("arn-2") is handler_after # noqa: SLF001 + assert invoker.unmarked_client_for("arn-2") is unmarked_after + for client in (unmarked_before, unmarked_after): + client.meta.events.register.assert_not_called() + + +def test_lambda_invoker_keeps_a_running_execution_and_its_targets_on_its_endpoint(): + """An endpoint update applies to executions not yet pinned. An + execution already invoked stays on its endpoint, and so do the + chained targets it dispatches afterwards.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.side_effect = lambda *_a, **_k: Mock() + invoker = LambdaInvoker.create("http://localhost:3001", "us-west-2") + pinned_handler = invoker._get_client_for_execution("running") # noqa: SLF001 + + invoker.update_endpoint("http://localhost:3002", "us-west-2") + + assert invoker._get_client_for_execution("running") is pinned_handler # noqa: SLF001 + running_unmarked = invoker.unmarked_client_for("running") + new_handler = invoker._get_client_for_execution("started-later") # noqa: SLF001 + new_unmarked = invoker.unmarked_client_for("started-later") + assert new_handler is not pinned_handler + assert new_unmarked is not running_unmarked + assert _clients_created(mock_boto3) == [ + ("http://localhost:3001", "us-west-2"), # handler client at create + ("http://localhost:3002", "us-west-2"), # handler client at update + ("http://localhost:3001", "us-west-2"), # unmarked, running execution + ("http://localhost:3002", "us-west-2"), # unmarked, later execution + ] + + +def test_lambda_invoker_propagates_read_timeout_to_per_endpoint_clients(): + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.return_value = Mock() + invoker = LambdaInvoker.create( + "http://localhost:3001", "us-west-2", read_timeout_seconds=5460 ) + invoker.update_endpoint("http://localhost:3002", "us-west-2") + invoker._get_client_for_execution("arn", "http://localhost:3003") # noqa: SLF001 + timeouts = [ + call.kwargs["config"].read_timeout + for call in mock_boto3.client.call_args_list + ] + assert timeouts == [5460, 5460, 5460] + + +def test_lambda_invoker_child_inherits_the_endpoint_its_parent_is_pinned_to(): + """A parent pinned to A, then an endpoint update to B: a child the + parent starts afterwards is invoked at A, and its own chained + targets go to A, while an unrelated execution started later goes to + B.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.side_effect = lambda *_a, **_k: Mock() + invoker = LambdaInvoker.create("http://a", "us-west-2") + parent_handler = invoker._get_client_for_execution("parent") # noqa: SLF001 + + invoker.update_endpoint("http://b", "us-west-2") + invoker.inherit_endpoint("child", "parent") + + assert invoker._get_client_for_execution("child") is parent_handler # noqa: SLF001 + assert invoker.unmarked_client_for("child") is invoker.unmarked_client_for( + "parent" + ) + assert invoker._get_client_for_execution("later") is invoker.lambda_client # noqa: SLF001 + assert _clients_created(mock_boto3) == [ + ("http://a", "us-west-2"), # handler client at create + ("http://b", "us-west-2"), # handler client at update + ("http://a", "us-west-2"), # unmarked client, parent and child + ] + # A grandchild inherits the same pin through the child. + invoker.inherit_endpoint("grandchild", "child") + assert invoker._get_client_for_execution("grandchild") is parent_handler # noqa: SLF001 + + +def test_lambda_invoker_per_execution_endpoint_is_not_where_chained_work_goes(): + """An execution's own endpoint serves its function alone, so its + chained targets and children go to the name-routing endpoint current + at its first chained dispatch, and stay there across an update.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.side_effect = lambda *_a, **_k: Mock() + invoker = LambdaInvoker.create("http://global", "us-west-2") + global_handler = invoker.lambda_client + own_handler = invoker._get_client_for_execution( # noqa: SLF001 + "parent", "http://own", "us-west-2" + ) + assert own_handler is not global_handler + + plain_target = invoker.unmarked_client_for("parent") + invoker.inherit_endpoint("child", "parent") + invoker.update_endpoint("http://later", "us-west-2") + + # The parent's handler still goes to its own endpoint; its plain + # targets and its child go to the endpoint pinned at first dispatch. + assert ( + invoker._get_client_for_execution("parent", "http://own", "us-west-2") # noqa: SLF001 + is own_handler + ) + assert invoker.unmarked_client_for("parent") is plain_target + assert invoker.unmarked_client_for("child") is plain_target + assert invoker._get_client_for_execution("child") is global_handler # noqa: SLF001 + assert invoker._get_client_for_execution("later") is invoker.lambda_client # noqa: SLF001 + assert _clients_created(mock_boto3) == [ + ("http://global", "us-west-2"), # handler client at create + ("http://own", "us-west-2"), # the parent's own endpoint, its region + ("http://global", "us-west-2"), # unmarked client, parent and child + ("http://later", "us-west-2"), # handler client at update + ] + + +def test_lambda_invoker_signs_an_own_endpoint_in_the_execution_region(): + """The signing region for an execution's own endpoint is the + execution's region, so an endpoint update to another region does not + change it. Without a region the historical default applies.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + mock_boto3.client.side_effect = lambda *_a, **_k: Mock() + invoker = LambdaInvoker.create("http://global", "us-west-2") + own_handler = invoker._get_client_for_execution( # noqa: SLF001 + "running", "http://own", "us-west-2" + ) + + invoker.update_endpoint("http://global", "eu-west-1") + + assert ( + invoker._get_client_for_execution("running", "http://own", "us-west-2") # noqa: SLF001 + is own_handler + ) + invoker._get_client_for_execution("no-region", "http://own") # noqa: SLF001 + assert _clients_created(mock_boto3) == [ + ("http://global", "us-west-2"), # handler client at create + ("http://own", "us-west-2"), # the running execution's own endpoint + ("http://global", "eu-west-1"), # handler client at update + ("http://own", "us-east-1"), # no region given: historical default + ] + + +def test_lambda_invoker_invoke_uses_the_region_given_for_an_own_endpoint(): + """invoke() hands the execution's region to the client lookup.""" + with patch("aws_durable_execution_sdk_python_testing.invoker.boto3") as mock_boto3: + + def new_client(*_a, **_k): + client = Mock() + payload = Mock() + payload.read.return_value = json.dumps({"Status": "SUCCEEDED"}).encode() + client.invoke.return_value = { + "StatusCode": 200, + "Payload": payload, + "ResponseMetadata": {"HTTPHeaders": {"x-amzn-RequestId": "r"}}, + } + return client + + mock_boto3.client.side_effect = new_client + invoker = LambdaInvoker.create("http://global", "us-west-2") + input_data = DurableExecutionInvocationInput( + durable_execution_arn="arn", + checkpoint_token="token", # noqa: S106 + initial_execution_state=InitialExecutionState( + operations=[], next_marker="" + ), + ) + + invoker.invoke( + "fn", input_data, endpoint_url="http://own", region_name="eu-west-1" + ) + + assert _clients_created(mock_boto3) == [ + ("http://global", "us-west-2"), + ("http://own", "eu-west-1"), + ] + own_client = invoker._endpoint_clients[("http://own", "eu-west-1")] # noqa: SLF001 + own_client.invoke.assert_called_once() + invoker.lambda_client.invoke.assert_not_called() + + +def test_in_process_invoker_has_no_endpoint_to_inherit(): + invoker = InProcessInvoker(Mock(), Mock()) + invoker.update_endpoint("http://ignored", "us-west-2") + invoker.inherit_endpoint("child", "parent") + + +class _RawBody(io.BytesIO): + """Minimal raw HTTP body for a fabricated botocore response.""" + + def stream(self, *_args, **_kwargs): + yield self.getvalue() + + +def _invoke_headers(monkeypatch, **client_kwargs) -> dict: + """Create a client, send one Invoke, and return the HTTP headers it put on the wire. + + Dummy credentials are set before the client is created because boto3 + resolves them at creation time. The request is answered at + ``before-send``, after botocore has built and signed it, so every + ``before-call`` hook (including the client's own) has already run. + """ + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test") + client = create_lambda_client("http://localhost:3001", "us-west-2", **client_kwargs) + captured: dict = {} + + def answer(request, **_kwargs): + # botocore stores prepared header values as bytes. + captured.update( + { + k: (v.decode() if isinstance(v, bytes) else v) + for k, v in request.headers.items() + } + ) + return AWSResponse( + request.url, 200, {"content-type": "application/json"}, _RawBody(b"{}") + ) + + client.meta.events.register("before-send.lambda.Invoke", answer) + client.invoke(FunctionName="target", InvocationType="RequestResponse", Payload="{}") + return captured + + +def test_handler_invocation_client_marks_its_invokes(monkeypatch): + """Every handler invocation identifies itself to the endpoint.""" + headers = _invoke_headers(monkeypatch) + assert headers[INVOCATION_MARKER_HEADER] == INVOCATION_MARKER_VALUE + assert INVOCATION_MARKER_HEADER == "X-Dex-Handler-Invoke" + assert INVOCATION_MARKER_VALUE == "true" + + +def test_unmarked_client_sends_plain_invokes(monkeypatch): + """A chained invoke of a non-durable target looks like any caller's Invoke.""" + assert INVOCATION_MARKER_HEADER not in _invoke_headers( + monkeypatch, mark_invocations=False + ) def test_lambda_invoker_create_invocation_input(): diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py index 10076c0e6..746f083d6 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/model_test.py @@ -70,6 +70,7 @@ WaitStartedDetails, WaitSucceededDetails, events_to_operations, + executed_version, ) @@ -3665,3 +3666,49 @@ def test_invocation_completed_details_from_json_dict_invalid_timestamp(): match="StartTimestamp and EndTimestamp cannot be null", ): InvocationCompletedDetails.from_json_dict(json_dict) + + +@pytest.mark.parametrize( + ("qualifier", "version"), + [ + (None, "$LATEST"), + ("$LATEST", "$LATEST"), + ("7", "7"), + ("prod", "$LATEST"), + ("$LATEST.PUBLISHED", "$LATEST.PUBLISHED"), + ], +) +def test_executed_version_reports_what_lambda_would_run(qualifier, version): + """A numeric qualifier and $LATEST.PUBLISHED are reported as given; + anything else runs $LATEST here, since the runner keeps no versions or + aliases.""" + assert executed_version(qualifier) == version + + +def test_execution_summary_reports_the_execution_identity(): + """The summary's FunctionArn is the execution's own region, account and + qualified name, not a fixed placeholder.""" + from aws_durable_execution_sdk_python_testing.execution import ( + Execution as StoredExecution, + ) + + stored = StoredExecution.new( + StartDurableExecutionInput( + account_id="210987654321", + function_name="orders", + function_qualifier="3", + execution_name="run-1", + execution_timeout_seconds=60, + execution_retention_period_days=1, + invocation_id="inv-1", + ) + ) + stored.region = "eu-west-1" + stored.start() + + summary = Execution.from_execution(stored, "RUNNING", "us-west-2") + + assert ( + summary.function_arn + == "arn:aws:lambda:eu-west-1:210987654321:function:orders:3" + ) diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/observer_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/observer_test.py index c5be5c4d2..56ce15f9c 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/observer_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/observer_test.py @@ -27,6 +27,7 @@ def __init__(self): self.on_timed_out_calls = [] self.on_stopped_calls = [] self.on_callback_created_calls = [] + self.on_chained_invoke_started_calls = [] def on_completed(self, execution_arn: str, result: str | None = None) -> None: self.on_completed_calls.append((execution_arn, result)) @@ -51,6 +52,18 @@ def on_callback_created( (execution_arn, operation_id, callback_options, callback_token) ) + def on_chained_invoke_started( + self, + execution_arn: str, + operation_id: str, + function_name: str, + tenant_id: str | None, + payload: str | None, + ) -> None: + self.on_chained_invoke_started_calls.append( + (execution_arn, operation_id, function_name, tenant_id, payload) + ) + # region ExecutionNotifier collects effects diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py index c6ae8fbe9..ab75f0d24 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/runner_test.py @@ -750,7 +750,10 @@ def test_durable_function_test_runner_init( mock_processor.assert_called_once() mock_client.assert_called_once() mock_invoker.assert_called_once_with( - handler, mock_client.return_value, max_page_bytes=5 * 1024 * 1024 + handler, + mock_client.return_value, + max_page_bytes=5 * 1024 * 1024, + region="us-west-2", ) mock_executor.assert_called_once() @@ -1031,7 +1034,10 @@ def test_durable_context_test_runner_init( mock_processor.assert_called_once() mock_client.assert_called_once() mock_invoker.assert_called_once_with( - decorated_handler, mock_client.return_value, max_page_bytes=5 * 1024 * 1024 + decorated_handler, + mock_client.return_value, + max_page_bytes=5 * 1024 * 1024, + region="us-west-2", ) mock_executor.assert_called_once() @@ -1084,7 +1090,10 @@ def test_durable_child_context_test_runner_init_with_args( mock_processor.assert_called_once() mock_client.assert_called_once() mock_invoker.assert_called_once_with( - decorated_handler, mock_client.return_value, max_page_bytes=5 * 1024 * 1024 + decorated_handler, + mock_client.return_value, + max_page_bytes=5 * 1024 * 1024, + region="us-west-2", ) mock_executor.assert_called_once() diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/runner_web_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/runner_web_test.py index a7cf8f4ff..9375b20b1 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/runner_web_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/runner_web_test.py @@ -12,11 +12,20 @@ from aws_durable_execution_sdk_python_testing.exceptions import ( DurableFunctionsLocalRunnerError, ) -from aws_durable_execution_sdk_python_testing.invoker import _LAMBDA_CLIENT_CONFIG +from aws_durable_execution_sdk_python_testing.child_dispatcher import ( + EndpointChildDispatcher, + FunctionConfigs, + UnconfiguredChildDispatcher, +) +from aws_durable_execution_sdk_python_testing.invoker import ( + DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + read_timeout_for, +) from aws_durable_execution_sdk_python_testing.runner import ( WebRunner, WebRunnerConfig, ) +from aws_durable_execution_sdk_python_testing.stores.base import StoreType from aws_durable_execution_sdk_python_testing.web.server import WebServiceConfig @@ -416,12 +425,18 @@ def test_should_handle_boto3_client_creation_with_custom_config(): # Act - Test public behavior runner.start() - # Assert - Verify boto3 client was called with correct parameters - mock_boto3_client.assert_called_once_with( - "lambda", - endpoint_url="http://custom-endpoint:8080", - region_name="eu-west-1", - config=_LAMBDA_CLIENT_CONFIG, + # Assert - Without a chained-invoke configuration only the + # handler-invocation client exists; it targets the configured + # endpoint and region. + assert len(mock_boto3_client.call_args_list) == 1 + kwargs = mock_boto3_client.call_args.kwargs + assert mock_boto3_client.call_args.args == ("lambda",) + assert kwargs["endpoint_url"] == "http://custom-endpoint:8080" + assert kwargs["region_name"] == "eu-west-1" + assert kwargs["config"].read_timeout == DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS + assert isinstance( + runner._executor._child_dispatcher, # noqa: SLF001 + UnconfiguredChildDispatcher, ) # Verify public behavior works @@ -442,18 +457,55 @@ def test_should_handle_boto3_client_creation_with_defaults(): # Act - Test public behavior runner.start() - # Assert - Verify boto3 client was called with default parameters - mock_boto3_client.assert_called_once_with( - "lambda", - endpoint_url="http://127.0.0.1:3001", # Default lambda_endpoint value - region_name="us-west-2", # Default value - config=_LAMBDA_CLIENT_CONFIG, - ) + # Assert - The client uses the default endpoint and region + assert len(mock_boto3_client.call_args_list) == 1 + kwargs = mock_boto3_client.call_args.kwargs + assert kwargs["endpoint_url"] == "http://127.0.0.1:3001" # Default value + assert kwargs["region_name"] == "us-west-2" # Default value # Verify public behavior works runner.stop() +def test_should_dispatch_plain_targets_with_the_invoker_unmarked_client(): + """With function configurations, non-durable chained targets use the + invoker's unmarked client for the current endpoint, and the read + timeout follows the emulated function timeout.""" + runner_config = WebRunnerConfig( + web_service=WebServiceConfig(host="localhost", port=5000), + function_configs=FunctionConfigs.from_value( + '{"ProcessPayment": {"DurableConfig": {}}, "LookupPrice": {}}' + ), + invocation_timeout_seconds=5400, + ) + runner = WebRunner(runner_config) + + with patch("boto3.client") as mock_boto3_client: + handler_client, dispatch_client = Mock(), Mock() + mock_boto3_client.side_effect = [handler_client, dispatch_client] + + runner.start() + + dispatcher = runner._executor._child_dispatcher # noqa: SLF001 + assert isinstance(dispatcher, EndpointChildDispatcher) + configs = dispatcher._function_configs.by_name # noqa: SLF001 + assert configs["ProcessPayment"].is_durable + assert not configs["LookupPrice"].is_durable + # The dispatch client is created on first use, for the same endpoint + # and region as the handler client, with the same read timeout. + assert dispatcher._client_provider("arn:parent") is dispatch_client # noqa: SLF001 + assert len(mock_boto3_client.call_args_list) == 2 + for c in mock_boto3_client.call_args_list: + assert c.kwargs["endpoint_url"] == runner_config.lambda_endpoint + assert c.kwargs["region_name"] == runner_config.local_runner_region + assert c.kwargs["config"].read_timeout == read_timeout_for(5400) == 5460 + # The handler client marks its invokes; the dispatch client does not. + handler_client.meta.events.register.assert_called_once() + dispatch_client.meta.events.register.assert_not_called() + + runner.stop() + + def test_should_propagate_boto3_client_creation_exceptions(): """Test that start() propagates boto3 client creation exceptions through public API.""" # Arrange @@ -484,8 +536,9 @@ def test_should_create_boto3_client_during_start(): # Act - Test public behavior runner.start() - # Assert - Verify boto3 client was created - mock_boto3_client.assert_called_once() + # Assert - One client, for handler invocations. (A dispatch client + # exists only with a chained-invoke configuration.) + assert mock_boto3_client.call_count == 1 # Verify public behavior works runner.stop() @@ -682,7 +735,11 @@ def test_should_create_all_required_dependencies_during_start(): mock_store_class.assert_called_once() mock_scheduler_class.assert_called_once() mock_invoker_class.assert_called_once_with( - mock_client, max_page_bytes=5 * 1024 * 1024 + mock_client, + max_page_bytes=5 * 1024 * 1024, + read_timeout_seconds=DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + endpoint_url=runner_config.lambda_endpoint, + region_name=runner_config.local_runner_region, ) # Verify Executor was called with the expected parameters including checkpoint_processor assert mock_executor_class.call_count == 1 @@ -768,17 +825,19 @@ def test_should_pass_correct_boto3_client_to_lambda_invoker(): # Act runner.start() - # Assert - Verify boto3 client was created with correct parameters - mock_boto3_client.assert_called_once_with( - "lambda", - endpoint_url="http://test-endpoint:7777", - region_name="ap-southeast-2", - config=_LAMBDA_CLIENT_CONFIG, - ) - - # Verify LambdaInvoker was created with the client + # Assert - The handler-invocation client is the one handed to + # LambdaInvoker, along with the read timeout it must apply to any + # per-endpoint client it creates later. + kwargs = mock_boto3_client.call_args.kwargs + assert mock_boto3_client.call_args.args == ("lambda",) + assert kwargs["endpoint_url"] == "http://test-endpoint:7777" + assert kwargs["region_name"] == "ap-southeast-2" mock_invoker_class.assert_called_once_with( - mock_client, max_page_bytes=5 * 1024 * 1024 + mock_client, + max_page_bytes=5 * 1024 * 1024, + read_timeout_seconds=DEFAULT_LAMBDA_READ_TIMEOUT_SECONDS, + endpoint_url=runner_config.lambda_endpoint, + region_name=runner_config.local_runner_region, ) # Cleanup @@ -1755,3 +1814,24 @@ def test_state_transitions_prevent_invalid_operations(): DurableFunctionsLocalRunnerError, match="Server not started" ): runner.serve_forever() + + +def test_web_runner_config_positional_construction_is_unchanged(): + """function_configs is the last field, so a caller that built the + config positionally before it existed gets the same values.""" + config = WebRunnerConfig( + WebServiceConfig(host="h", port=1), + "http://lambda:3001", + "http://runner:5000", + "eu-west-1", + "local", + StoreType.MEMORY, + None, + 60, + True, + 1024, + ) + + assert config.skip_time is True + assert config.max_invocation_page_bytes == 1024 + assert config.function_configs is None diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/threads_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/threads_test.py new file mode 100644 index 000000000..cfb36d84b --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/threads_test.py @@ -0,0 +1,134 @@ +"""Tests for the bounded daemon-thread pool.""" + +from __future__ import annotations + +import threading +import time +from concurrent.futures import CancelledError + +import pytest + +from aws_durable_execution_sdk_python_testing.threads import ( + DaemonThreadPool, + default_max_workers, +) + + +def test_work_runs_on_a_daemon_thread(): + pool = DaemonThreadPool(thread_name_prefix="t") + + future = pool.submit(lambda: threading.current_thread()) + worker = future.result(timeout=5) + + assert worker.daemon is True + assert worker.name.startswith("t_") + pool.shutdown() + + +def test_default_cap_matches_thread_pool_executor(): + assert DaemonThreadPool().max_workers == default_max_workers() + with pytest.raises(ValueError, match="greater than 0"): + DaemonThreadPool(max_workers=0) + + +def test_queued_work_runs_in_submission_order_beyond_the_cap(): + pool = DaemonThreadPool(max_workers=1) + gate = threading.Event() + order: list[int] = [] + + def first() -> None: + gate.wait(5) + order.append(0) + + futures = [pool.submit(first)] + [pool.submit(order.append, n) for n in (1, 2, 3)] + gate.set() + for future in futures: + future.result(timeout=5) + + assert order == [0, 1, 2, 3] + pool.shutdown() + + +def test_never_runs_more_than_the_cap_at_once(): + pool = DaemonThreadPool(max_workers=2) + release = threading.Event() + running = 0 + peak = 0 + lock = threading.Lock() + + def work() -> None: + nonlocal running, peak + with lock: + running += 1 + peak = max(peak, running) + release.wait(5) + with lock: + running -= 1 + + futures = [pool.submit(work) for _ in range(6)] + time.sleep(0.2) + with lock: + assert running == 2 + release.set() + for future in futures: + future.result(timeout=5) + + assert peak == 2 + pool.shutdown() + + +def test_an_idle_worker_is_reused_instead_of_starting_another(): + pool = DaemonThreadPool(max_workers=4) + + pool.submit(lambda: None).result(timeout=5) + pool.submit(lambda: None).result(timeout=5) + + assert len(pool._threads) == 1 # noqa: SLF001 + pool.shutdown() + + +def test_shutdown_cancels_queued_work_and_refuses_new_work(): + pool = DaemonThreadPool(max_workers=1) + release = threading.Event() + busy = pool.submit(release.wait, 5) + queued = pool.submit(lambda: "never") + + pool.shutdown() + + assert queued.cancelled() + with pytest.raises(CancelledError): + queued.result(timeout=0) + with pytest.raises(RuntimeError, match="after shutdown"): + pool.submit(lambda: None) + release.set() + assert busy.result(timeout=5) is True + pool.shutdown() # idempotent + + +def test_shutdown_returns_while_a_worker_is_blocked(): + """The blocked worker is left alone: shutdown neither waits for it + nor fails. Its result lands on a future nobody reads.""" + pool = DaemonThreadPool(max_workers=1) + release = threading.Event() + blocked = pool.submit(release.wait, 30) + + started = time.monotonic() + pool.shutdown(wait=True) + elapsed = time.monotonic() - started + + assert elapsed < 1 + assert not blocked.done() + release.set() + assert blocked.result(timeout=5) is True + + +def test_an_exception_lands_on_the_future(): + pool = DaemonThreadPool(max_workers=1) + + def boom() -> None: + msg = "boom" + raise ValueError(msg) + + with pytest.raises(ValueError, match="boom"): + pool.submit(boom).result(timeout=5) + pool.shutdown() diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/close_while_invoke_blocked_int_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/close_while_invoke_blocked_int_test.py new file mode 100644 index 000000000..8de035884 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/close_while_invoke_blocked_int_test.py @@ -0,0 +1,106 @@ +"""Process exit while a handler invocation is blocked. + +Python cannot interrupt a thread blocked in a socket read, and the +interpreter joins non-daemon threads at exit. So the proof that a +blocked Invoke cannot hold the runner's process is a process: a real +WebRunner whose Lambda endpoint accepts the connection and never +answers, stopped while that invocation is in flight. The child process +must exit long before the read timeout (the invocation timeout plus +sixty seconds). +""" + +from __future__ import annotations + +import os +import subprocess +import sys +import textwrap +import time + +_SCRIPT = textwrap.dedent( + """ + import socket + import threading + import time + + from aws_durable_execution_sdk_python_testing.model import ( + StartDurableExecutionInput, + ) + from aws_durable_execution_sdk_python_testing.runner import ( + WebRunner, + WebRunnerConfig, + ) + from aws_durable_execution_sdk_python_testing.web.server import WebServiceConfig + + # A Lambda endpoint that accepts every connection and never answers. + accepted = threading.Event() + endpoint = socket.socket() + endpoint.bind(("127.0.0.1", 0)) + endpoint.listen() + held: list[socket.socket] = [] + + def accept_forever() -> None: + while True: + conn, _ = endpoint.accept() + held.append(conn) + accepted.set() + + threading.Thread(target=accept_forever, daemon=True).start() + + probe = socket.socket() + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + probe.close() + + runner = WebRunner( + WebRunnerConfig( + web_service=WebServiceConfig(host="127.0.0.1", port=port), + lambda_endpoint=f"http://127.0.0.1:{endpoint.getsockname()[1]}", + local_runner_endpoint=f"http://127.0.0.1:{port}", + ) + ) + runner.start() + runner._executor.start_execution( # noqa: SLF001 + StartDurableExecutionInput( + account_id="123456789012", + function_name="blocked", + function_qualifier="$LATEST", + execution_name="blocked-run", + execution_timeout_seconds=300, + execution_retention_period_days=1, + ) + ) + if not accepted.wait(15): + raise SystemExit("the handler invocation never reached the endpoint") + time.sleep(0.3) # let the invoke settle into its blocked read + print("invoke in flight", flush=True) + runner.stop() + print("runner stopped", flush=True) + """ +) + + +def test_stopping_the_runner_during_a_blocked_invoke_lets_the_process_exit(): + env = dict(os.environ) + # The runner signs its Invokes; any credentials will do locally. + env.setdefault("AWS_ACCESS_KEY_ID", "test") + env.setdefault("AWS_SECRET_ACCESS_KEY", "test") + env.setdefault("AWS_DEFAULT_REGION", "us-west-2") + + started = time.monotonic() + completed = subprocess.run( # noqa: S603 + [sys.executable, "-c", _SCRIPT], + check=False, + capture_output=True, + text=True, + timeout=60, + env=env, + ) + elapsed = time.monotonic() - started + + assert completed.returncode == 0, completed.stderr + assert "invoke in flight" in completed.stdout + assert "runner stopped" in completed.stdout + # Import, start, invoke and stop take about a second. The blocked + # read would hold a non-daemon thread for 960 seconds. + assert elapsed < 15, f"process took {elapsed:.1f}s to exit" diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/routes_arn_encoding_int_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/routes_arn_encoding_int_test.py index 9ddbbf7df..9c5dc7a3f 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/routes_arn_encoding_int_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/web/e2e/routes_arn_encoding_int_test.py @@ -52,6 +52,9 @@ def invoke(self, *args: Any, **kwargs: Any) -> Any: # noqa: ARG002 def update_endpoint(self, *args: Any, **kwargs: Any) -> None: # noqa: ARG002 return None + def inherit_endpoint(self, *args: Any, **kwargs: Any) -> None: # noqa: ARG002 + return None + def _assert_no_percent_encoding_in_error(exc: ClientError, arn: str) -> None: """Fail the test if a ResourceNotFoundException carries a %2F-form ARN. diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/web/handlers_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/web/handlers_test.py index b56bf2767..ae5ad8c51 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/web/handlers_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/web/handlers_test.py @@ -2269,6 +2269,7 @@ def test_update_lambda_endpoint_handler_success(): ) executor = Mock() + executor.region = "us-west-2" lambda_invoker = Mock(spec=LambdaInvoker) executor._invoker = lambda_invoker # noqa: SLF001 handler = UpdateLambdaEndpointHandler(executor) @@ -2293,6 +2294,45 @@ def test_update_lambda_endpoint_handler_success(): ) +def test_update_lambda_endpoint_handler_rejects_another_region(): + """The runner emulates one region, fixed at startup. A request that + names another region is refused, so the invoker and the executions + never disagree on the region.""" + from aws_durable_execution_sdk_python_testing.invoker import LambdaInvoker + from aws_durable_execution_sdk_python_testing.web.handlers import ( + UpdateLambdaEndpointHandler, + ) + from aws_durable_execution_sdk_python_testing.web.routes import ( + UpdateLambdaEndpointRoute, + ) + + executor = Mock() + executor.region = "us-west-2" + lambda_invoker = Mock(spec=LambdaInvoker) + executor._invoker = lambda_invoker # noqa: SLF001 + handler = UpdateLambdaEndpointHandler(executor) + + base_route = Route.from_string("/lambda-endpoint") + update_route = UpdateLambdaEndpointRoute.from_route(base_route) + + request = HTTPRequest( + method="PUT", + path=update_route, + headers={"Content-Type": "application/json"}, + query_params={}, + body={"EndpointUrl": "http://localhost:8080", "RegionName": "eu-west-1"}, + ) + + response = handler.handle(update_route, request) + + assert response.status_code == 400 + assert response.body["Type"] == "InvalidParameterValueException" + assert response.body["message"] == ( + "RegionName must be us-west-2, the region this runner emulates; got eu-west-1" + ) + lambda_invoker.update_endpoint.assert_not_called() + + def test_update_lambda_endpoint_handler_missing_endpoint_url(): """Test UpdateLambdaEndpointHandler with missing EndpointUrl.""" from aws_durable_execution_sdk_python_testing.web.handlers import ( @@ -2324,7 +2364,7 @@ def test_update_lambda_endpoint_handler_missing_endpoint_url(): def test_update_lambda_endpoint_handler_default_region(): - """Test UpdateLambdaEndpointHandler uses default region when not specified.""" + """Without RegionName the endpoint moves and the runner's region stays.""" from aws_durable_execution_sdk_python_testing.invoker import LambdaInvoker from aws_durable_execution_sdk_python_testing.web.handlers import ( UpdateLambdaEndpointHandler, @@ -2334,6 +2374,7 @@ def test_update_lambda_endpoint_handler_default_region(): ) executor = Mock() + executor.region = "eu-central-1" lambda_invoker = Mock(spec=LambdaInvoker) executor._invoker = lambda_invoker # noqa: SLF001 handler = UpdateLambdaEndpointHandler(executor) @@ -2353,7 +2394,7 @@ def test_update_lambda_endpoint_handler_default_region(): assert response.status_code == 200 lambda_invoker.update_endpoint.assert_called_once_with( - "http://localhost:8080", "us-east-1" + "http://localhost:8080", "eu-central-1" ) diff --git a/packages/aws-durable-execution-sdk-python-testing/tests/worker/registry_test.py b/packages/aws-durable-execution-sdk-python-testing/tests/worker/registry_test.py index 0b50cc0c4..405246628 100644 --- a/packages/aws-durable-execution-sdk-python-testing/tests/worker/registry_test.py +++ b/packages/aws-durable-execution-sdk-python-testing/tests/worker/registry_test.py @@ -5,7 +5,11 @@ from typing import TYPE_CHECKING, cast import pytest +from unittest.mock import Mock +from aws_durable_execution_sdk_python_testing.stores.memory import ( + InMemoryExecutionStore, +) from aws_durable_execution_sdk_python_testing.worker.registry import ExecutionRegistry from aws_durable_execution_sdk_python_testing.worker.task import ( ExecutionTask, @@ -84,3 +88,19 @@ def test_shutdown_stops_all_lanes_and_drops_workers() -> None: for worker in workers: with pytest.raises(RuntimeError): worker.submit(_EchoTask("x")) + + +def test_registry_refuses_work_after_shutdown(): + """A lane created after shutdown would outlive the runner. So submit + and get_or_create raise instead of creating one.""" + from aws_durable_execution_sdk_python_testing.worker.checkpoint_tasks import ( + CallableTask, + ) + + registry = ExecutionRegistry(InMemoryExecutionStore(), Mock()) + registry.shutdown() + + with pytest.raises(RuntimeError, match="shut down"): + registry.get_or_create("arn-1") + with pytest.raises(RuntimeError, match="shut down"): + registry.submit("arn-1", CallableTask(lambda: None))