Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
50 changes: 50 additions & 0 deletions packages/aws-durable-execution-sdk-python-testing/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
Original file line number Diff line number Diff line change
@@ -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
),
)
Comment thread
yaythomas marked this conversation as resolved.

notifier.notify_chained_invoke_started(
Comment thread
yaythomas marked this conversation as resolved.
Comment thread
yaythomas marked this conversation as resolved.
execution_arn=execution_arn,
operation_id=update.operation_id,
function_name=options.function_name,
tenant_id=options.tenant_id,
payload=update.payload,
)
Comment thread
yaythomas marked this conversation as resolved.
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
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand All @@ -40,6 +43,7 @@
from datetime import datetime

from aws_durable_execution_sdk_python.lambda_service import (
ErrorObject,
OperationUpdate,
)

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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:
Expand All @@ -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

Expand Down
Loading
Loading