From 1ae73e8bf2dc59b81c18818e8e75b75c30f70cc7 Mon Sep 17 00:00:00 2001 From: nvasiu Date: Tue, 22 Sep 2026 20:39:27 +0000 Subject: [PATCH] feat(dmap): add distributed map operation Add ctx.distributed_map with inline, S3, and reader sources, the config, processor, completion, and destination types, the result types, and function-authoring helpers for item and batch handlers. Serialization lives in lambda_service.py with one dataclass per API shape. Translation happens only in the executor, so config.py has no runtime import from lambda_service. Tests are split by layer. Items, Output, and the record body are opaque strings, matching the API model, so any serdes works rather than JSON only. Every terminal state resolves with the summary, and throw_if_error() opts into raising. Missing details or a missing completion reason raise ExecutionError. The processor factories are batch, item_failures, and item_results. Source and destination factories are InlineSource, S3Source, ReaderSource, and S3Destination. Retry is max_retry_attempts and max_retry_duration on the processor. Passing DistributedMapResultConfig selects the DistributedMapResult return type statically. Only the result types and their enums are exported from the package root. --- .../__init__.py | 30 + .../config.py | 669 +++++++- .../context.py | 82 + .../dmap/__init__.py | 0 .../dmap/handlers.py | 453 ++++++ .../dmap/models.py | 127 ++ .../exceptions.py | 4 + .../lambda_service.py | 631 +++++++- .../operation/dmap.py | 582 +++++++ .../plugin.py | 1 + .../aws_durable_execution_sdk_python/state.py | 3 + .../tests/config_test.py | 171 ++ .../tests/context_test.py | 210 +++ .../tests/dmap/handlers_test.py | 253 +++ .../tests/dmap/models_test.py | 173 ++ .../tests/e2e/dmap_helpers_int_test.py | 146 ++ .../tests/e2e/dmap_int_test.py | 246 +++ .../tests/exceptions_test.py | 2 + .../tests/lambda_service_test.py | 96 ++ .../tests/operation/dmap_test.py | 1388 +++++++++++++++++ 20 files changed, 5264 insertions(+), 3 deletions(-) create mode 100644 packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/__init__.py create mode 100644 packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/handlers.py create mode 100644 packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/models.py create mode 100644 packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/dmap.py create mode 100644 packages/aws-durable-execution-sdk-python/tests/dmap/handlers_test.py create mode 100644 packages/aws-durable-execution-sdk-python/tests/dmap/models_test.py create mode 100644 packages/aws-durable-execution-sdk-python/tests/e2e/dmap_helpers_int_test.py create mode 100644 packages/aws-durable-execution-sdk-python/tests/e2e/dmap_int_test.py create mode 100644 packages/aws-durable-execution-sdk-python/tests/operation/dmap_test.py diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/__init__.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/__init__.py index 69f7bb16..bd019953 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/__init__.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/__init__.py @@ -16,6 +16,9 @@ CompletionItemStatus, CompletionOutcome, CompletionStatus, + DistributedMapCompletionReason, + DistributedMapItemStatus, + DistributedMapStatus, ParallelBranch, complete_batch, continue_batch, @@ -27,6 +30,19 @@ durable_wait_for_callback, durable_with_child_context, ) +from aws_durable_execution_sdk_python.dmap.handlers import ( + distributed_map_batch_handler, + distributed_map_item_handler, + distributed_map_reader, + durable_distributed_map_batch_handler, + durable_distributed_map_item_handler, +) +from aws_durable_execution_sdk_python.dmap.models import ( + DistributedMapItemError, + DistributedMapResult, + DistributedMapResultItem, + DistributedMapSummary, +) # Most common exceptions - users need to handle these exceptions from aws_durable_execution_sdk_python.exceptions import ( @@ -35,6 +51,7 @@ CallbackSubmitterError, CallbackTimeoutError, ChildContextError, + DistributedMapError, DurableExecutionsError, DurableOperationError, ExecutionError, @@ -69,6 +86,14 @@ "CompletionItemStatus", "CompletionOutcome", "CompletionStatus", + "DistributedMapCompletionReason", + "DistributedMapError", + "DistributedMapItemError", + "DistributedMapItemStatus", + "DistributedMapResult", + "DistributedMapResultItem", + "DistributedMapStatus", + "DistributedMapSummary", "DurableContext", "DurableExecutionsError", "DurableOperationError", @@ -87,6 +112,11 @@ "__version__", "complete_batch", "continue_batch", + "distributed_map_batch_handler", + "distributed_map_item_handler", + "distributed_map_reader", + "durable_distributed_map_batch_handler", + "durable_distributed_map_item_handler", "durable_execution", "durable_parallel_branch", "durable_step", diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/config.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/config.py index aa5a2b32..111af630 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/config.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/config.py @@ -6,17 +6,71 @@ import random from dataclasses import dataclass, field from enum import Enum, StrEnum -from typing import TYPE_CHECKING, Generic, TypeVar +from typing import TYPE_CHECKING, Any, ClassVar, Generic, Literal, TypeVar from aws_durable_execution_sdk_python.exceptions import ValidationError +class DistributedMapStatus(Enum): + SUCCEEDED = "SUCCEEDED" + FAILED = "FAILED" + STOPPED = "STOPPED" + TIMED_OUT = "TIMED_OUT" + + +class DistributedMapItemStatus(Enum): + """Terminal status of a single map run item.""" + + SUCCEEDED = "SUCCEEDED" + FAILED = "FAILED" + + +class DistributedMapCompletionReason(Enum): + ALL_COMPLETED = "ALL_COMPLETED" + ITEM_LIMIT_REACHED = "ITEM_LIMIT_REACHED" + STOPPED = "STOPPED" + TIMED_OUT = "TIMED_OUT" + FAILURE_TOLERANCE_EXCEEDED = "FAILURE_TOLERANCE_EXCEEDED" + SOURCE_FAILED = "SOURCE_FAILED" + DESTINATION_FAILED = "DESTINATION_FAILED" + INLINE_RESULT_LIMIT_EXCEEDED = "INLINE_RESULT_LIMIT_EXCEEDED" + INVALID_CONFIGURATION = "INVALID_CONFIGURATION" + QUOTA_EXCEEDED = "QUOTA_EXCEEDED" + KMS_ACCESS_DENIED = "KMS_ACCESS_DENIED" + INTERNAL_ERROR = "INTERNAL_ERROR" + UNKNOWN_TO_SDK_VERSION = "UNKNOWN_TO_SDK_VERSION" + + +class DistributedMapSourceFormat(Enum): + JSON_LINES = "JSON_LINES" + JSON_ARRAY = "JSON_ARRAY" + CSV = "CSV" + + +class ProcessorResponseMode(Enum): + """How a processor function reports per-item outcomes.""" + + BATCH = "BATCH" + ITEM_FAILURES = "ITEM_FAILURES" + ITEM_RESULTS = "ITEM_RESULTS" + + +class DistributedMapCsvDelimiter(Enum): + """Column delimiter for a CSV distributed map source.""" + + COMMA = "COMMA" + PIPE = "PIPE" + SEMICOLON = "SEMICOLON" + SPACE = "SPACE" + TAB = "TAB" + + P = TypeVar("P") # Payload type R = TypeVar("R") # Result type T = TypeVar("T") if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Sequence from aws_durable_execution_sdk_python.lambda_service import OperationSubType from aws_durable_execution_sdk_python.retries import RetryDecision @@ -490,6 +544,617 @@ class StepConfig: serdes: SerDes | None = None +# region map run configuration + + +@dataclass(frozen=True) +class S3Uri: + """A parsed ``s3://bucket/path`` URI.""" + + bucket: str + path: str | None = None + + @classmethod + def parse(cls, uri: str) -> S3Uri: + """Parse an ``s3://bucket/path`` URI, rejecting any other scheme.""" + if not uri.startswith("s3://"): + msg = f"S3 URI must start with s3://, got: {uri}" + raise ValidationError(msg) + bucket, _, path = uri.removeprefix("s3://").partition("/") + if not bucket: + msg = f"Invalid S3 URI: {uri}" + raise ValidationError(msg) + return cls(bucket=bucket, path=path or None) + + +def _validate_bucket_owner(value: str | None) -> None: + """Validate an expected bucket owner is a 12-digit account id.""" + if value is not None and ( + len(value) != 12 or not (value.isascii() and value.isdigit()) # noqa: PLR2004 + ): + msg = f"expected_bucket_owner must be a 12-digit account id, got: {value}" + raise ValidationError(msg) + + +def _validate_columns(name: str, columns: tuple[str, ...] | None) -> None: + """Validate a CSV columns/headers tuple is non-empty with no duplicates.""" + if columns is None: + return + if len(columns) == 0: + msg = f"{name} must be non-empty" + raise ValidationError(msg) + if len(set(columns)) != len(columns): + msg = f"{name} must not contain duplicates" + raise ValidationError(msg) + + +# NamespacedFunctionName max length; the service enforces the full grammar. +_MAX_FUNCTION_NAME_LENGTH = 170 + + +def _validate_function_name(name: str) -> None: + """Validate a Lambda function reference is present and within the length limit.""" + if not name: + msg = "function name must be non-empty" + raise ValidationError(msg) + if len(name) > _MAX_FUNCTION_NAME_LENGTH: + msg = f"function name must be at most {_MAX_FUNCTION_NAME_LENGTH} characters" + raise ValidationError(msg) + + +@dataclass(frozen=True) +class DistributedMapCompletionConfig: + """Failure-tolerance configuration for a map run.""" + + tolerated_failure_count: int | None = None + tolerated_failure_percentage: float | None = None + minimum_sample_size: int | None = None + + def __post_init__(self) -> None: + if ( + self.tolerated_failure_count is not None + and self.tolerated_failure_percentage is not None + ): + msg = ( + "tolerated_failure_count and tolerated_failure_percentage " + "are mutually exclusive" + ) + raise ValidationError(msg) + if ( + self.minimum_sample_size is not None + and self.tolerated_failure_percentage is None + ): + msg = "minimum_sample_size is only valid with tolerated_failure_percentage" + raise ValidationError(msg) + if ( + self.tolerated_failure_count is not None + and self.tolerated_failure_count < 0 + ): + msg = ( + "tolerated_failure_count must be non-negative, got: " + f"{self.tolerated_failure_count}" + ) + raise ValidationError(msg) + if self.tolerated_failure_percentage is not None and not ( + 0 <= self.tolerated_failure_percentage <= 100 # noqa: PLR2004 + ): + msg = ( + "tolerated_failure_percentage must be between 0 and 100, got: " + f"{self.tolerated_failure_percentage}" + ) + raise ValidationError(msg) + if self.minimum_sample_size is not None and self.minimum_sample_size < 1: + msg = ( + "minimum_sample_size must be at least 1, got: " + f"{self.minimum_sample_size}" + ) + raise ValidationError(msg) + + @staticmethod + def failure_count(count: int) -> DistributedMapCompletionConfig: + """Abort once this many items have permanently failed.""" + return DistributedMapCompletionConfig(tolerated_failure_count=count) + + @staticmethod + def failure_percentage( + percentage: float, *, minimum_sample_size: int | None = None + ) -> DistributedMapCompletionConfig: + """Abort once the failure rate exceeds this percentage.""" + return DistributedMapCompletionConfig( + tolerated_failure_percentage=percentage, + minimum_sample_size=minimum_sample_size, + ) + + +@dataclass(frozen=True) +class DistributedMapProcessor: + """Processor configuration for a map run.""" + + UNLIMITED: ClassVar[str] = "unlimited" + + function_name: str + response_mode: ProcessorResponseMode = ProcessorResponseMode.BATCH + batch_size: int | None = None + max_retry_attempts: int | Literal["unlimited"] | None = None + max_retry_duration: Duration | None = None + durable_execution_name_prefix: str | None = None + + def __post_init__(self) -> None: + _validate_function_name(self.function_name) + if self.batch_size is not None and not ( + 1 <= self.batch_size <= 10000 # noqa: PLR2004 + ): + msg = f"batch_size must be between 1 and 10000, got: {self.batch_size}" + raise ValidationError(msg) + if isinstance(self.max_retry_attempts, int) and self.max_retry_attempts < 0: + msg = ( + "max_retry_attempts must be non-negative or " + "DistributedMapProcessor.UNLIMITED, " + f"got: {self.max_retry_attempts}" + ) + raise ValidationError(msg) + if self.max_retry_duration is not None and not ( + 60 <= self.max_retry_duration.to_seconds() <= 21600 # noqa: PLR2004 + ): + msg = ( + "max_retry_duration must be between 1 minute and 6 hours, got: " + f"{self.max_retry_duration.to_seconds()}s" + ) + raise ValidationError(msg) + if self.durable_execution_name_prefix is not None and not ( + 1 <= len(self.durable_execution_name_prefix) <= 36 # noqa: PLR2004 + ): + msg = ( + "durable_execution_name_prefix must be 1 to 36 characters, got: " + f"{len(self.durable_execution_name_prefix)}" + ) + raise ValidationError(msg) + + @classmethod + def batch( + cls, + name: str, + *, + batch_size: int | None = None, + max_retry_attempts: int | Literal["unlimited"] | None = None, + max_retry_duration: Duration | None = None, + durable_execution_name_prefix: str | None = None, + ) -> DistributedMapProcessor: + """Processor that reports a single pass/fail outcome for the whole batch, with no per-item results.""" + return cls( + function_name=name, + response_mode=ProcessorResponseMode.BATCH, + batch_size=batch_size, + max_retry_attempts=max_retry_attempts, + max_retry_duration=max_retry_duration, + durable_execution_name_prefix=durable_execution_name_prefix, + ) + + @classmethod + def item_failures( + cls, + name: str, + *, + batch_size: int | None = None, + max_retry_attempts: int | Literal["unlimited"] | None = None, + max_retry_duration: Duration | None = None, + durable_execution_name_prefix: str | None = None, + ) -> DistributedMapProcessor: + """Processor that reports the ids of failed items, with all others marked succeeded.""" + return cls( + function_name=name, + response_mode=ProcessorResponseMode.ITEM_FAILURES, + batch_size=batch_size, + max_retry_attempts=max_retry_attempts, + max_retry_duration=max_retry_duration, + durable_execution_name_prefix=durable_execution_name_prefix, + ) + + @classmethod + def item_results( + cls, + name: str, + *, + batch_size: int | None = None, + max_retry_attempts: int | Literal["unlimited"] | None = None, + max_retry_duration: Duration | None = None, + durable_execution_name_prefix: str | None = None, + ) -> DistributedMapProcessor: + """Processor that reports the results (output or error) for every item.""" + return cls( + function_name=name, + response_mode=ProcessorResponseMode.ITEM_RESULTS, + batch_size=batch_size, + max_retry_attempts=max_retry_attempts, + max_retry_duration=max_retry_duration, + durable_execution_name_prefix=durable_execution_name_prefix, + ) + + +def _parse_delimiter( + value: str | DistributedMapCsvDelimiter, +) -> DistributedMapCsvDelimiter: + """Normalize a delimiter passed as a string or enum member.""" + if isinstance(value, DistributedMapCsvDelimiter): + return value + try: + return DistributedMapCsvDelimiter(value) + except ValueError: + allowed = ", ".join(d.value for d in DistributedMapCsvDelimiter) + msg = f"delimiter must be one of ({allowed}), got: {value!r}" + raise ValidationError(msg) from None + + +@dataclass(frozen=True) +class S3SourceConfig: + """Resolved S3 source configuration.""" + + bucket: str + key: str | None = None + prefix: str | None = None + fmt: DistributedMapSourceFormat | None = None + delimiter: DistributedMapCsvDelimiter | None = None + headers: tuple[str, ...] | None = None + expected_bucket_owner: str | None = None + + def __post_init__(self) -> None: + _validate_bucket_owner(self.expected_bucket_owner) + _validate_columns("headers", self.headers) + + +@dataclass(frozen=True) +class ReaderSourceConfig: + """Resolved reader-function source configuration.""" + + function_name: str + initial_state: Any = None + state_serdes: SerDes | None = None # None = DEFAULT_JSON_SERDES + + def __post_init__(self) -> None: + _validate_function_name(self.function_name) + + +@dataclass(frozen=True) +class DistributedMapSource: + """Source configuration for a map run.""" + + max_items: int | None = None + inline_items: tuple[Any, ...] | None = None + inline_serdes: SerDes | None = None # None = DEFAULT_JSON_SERDES + s3: S3SourceConfig | None = None + reader: ReaderSourceConfig | None = None + + def __post_init__(self) -> None: + if self.max_items is not None and self.max_items < 1: + msg = f"max_items must be at least 1, got: {self.max_items}" + raise ValidationError(msg) + populated = sum( + 1 + for value in (self.inline_items, self.s3, self.reader) + if value is not None + ) + if populated != 1: + msg = "exactly one of inline_items, s3 or reader must be set" + raise ValidationError(msg) + + +class InlineSource: + """Inline source factory.""" + + @staticmethod + def of( + items: Sequence[Any], + *, + serdes: SerDes | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """An in-memory list of items embedded in the start checkpoint.""" + return DistributedMapSource( + inline_items=tuple(items), + inline_serdes=serdes, + max_items=max_items, + ) + + +class S3Source: + """S3 source factories.""" + + @staticmethod + def json_lines( + uri: str, + *, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a single object, treating each line as an item.""" + parsed_uri = S3Uri.parse(uri) + bucket, key = parsed_uri.bucket, parsed_uri.path + if key is None: + msg = "json_lines requires an S3 object key" + raise ValidationError(msg) + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + key=key, + fmt=DistributedMapSourceFormat.JSON_LINES, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def json_array( + uri: str, + *, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a single object holding a JSON array, treating each element as an item.""" + parsed_uri = S3Uri.parse(uri) + bucket, key = parsed_uri.bucket, parsed_uri.path + if key is None: + msg = "json_array requires an S3 object key" + raise ValidationError(msg) + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + key=key, + fmt=DistributedMapSourceFormat.JSON_ARRAY, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def csv( + uri: str, + *, + headers: Sequence[str] | None = None, + delimiter: str | DistributedMapCsvDelimiter = DistributedMapCsvDelimiter.COMMA, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a single object, treating each record as an item.""" + parsed_uri = S3Uri.parse(uri) + bucket, key = parsed_uri.bucket, parsed_uri.path + if key is None: + msg = "csv requires an S3 object key" + raise ValidationError(msg) + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + key=key, + fmt=DistributedMapSourceFormat.CSV, + delimiter=_parse_delimiter(delimiter), + headers=tuple(headers) if headers is not None else None, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def objects( + prefix_uri: str, + *, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read each object under a prefix as one item.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + prefix=prefix or "", + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def flattened_json_lines( + prefix_uri: str, + *, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a prefix, flattening each object's lines into items.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + prefix=prefix or "", + fmt=DistributedMapSourceFormat.JSON_LINES, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def flattened_json_array( + prefix_uri: str, + *, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a prefix, flattening each object's JSON array elements into items.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + prefix=prefix or "", + fmt=DistributedMapSourceFormat.JSON_ARRAY, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + @staticmethod + def flattened_csv( + prefix_uri: str, + *, + headers: Sequence[str] | None = None, + delimiter: str | DistributedMapCsvDelimiter = DistributedMapCsvDelimiter.COMMA, + expected_bucket_owner: str | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Read a prefix, flattening each object's records into items.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return DistributedMapSource( + max_items=max_items, + s3=S3SourceConfig( + bucket=bucket, + prefix=prefix or "", + fmt=DistributedMapSourceFormat.CSV, + delimiter=_parse_delimiter(delimiter), + headers=tuple(headers) if headers is not None else None, + expected_bucket_owner=expected_bucket_owner, + ), + ) + + +class ReaderSource: + """Reader-function source factories.""" + + @staticmethod + def from_function( + name: str, + *, + initial_state: Any = None, + state_serdes: SerDes | None = None, + max_items: int | None = None, + ) -> DistributedMapSource: + """Page items from a customer-supplied reader Lambda function.""" + return DistributedMapSource( + max_items=max_items, + reader=ReaderSourceConfig( + function_name=name, + initial_state=initial_state, + state_serdes=state_serdes, + ), + ) + + +@dataclass(frozen=True) +class SuccessDestination: + """S3 destination for succeeded items.""" + + bucket: str + prefix: str + include_input: bool = False + include_output: bool = True + expected_bucket_owner: str | None = None + + def __post_init__(self) -> None: + _validate_bucket_owner(self.expected_bucket_owner) + if not (self.include_input or self.include_output): + msg = "success destination must include input or output" + raise ValidationError(msg) + + +@dataclass(frozen=True) +class FailureDestination: + """S3 destination for permanently-failed items.""" + + bucket: str + prefix: str + include_input: bool = True + include_error: bool = True + expected_bucket_owner: str | None = None + + def __post_init__(self) -> None: + _validate_bucket_owner(self.expected_bucket_owner) + if not (self.include_input or self.include_error): + msg = "failure destination must include input or error" + raise ValidationError(msg) + + +@dataclass(frozen=True) +class DistributedMapDestinationConfig: + """Destination routing for map run results.""" + + on_success: SuccessDestination | None = None + on_failure: FailureDestination | None = None + + +class S3Destination: + """S3 destination factories.""" + + @staticmethod + def successes( + prefix_uri: str, + *, + include_input: bool = False, + include_output: bool = True, + expected_bucket_owner: str | None = None, + ) -> SuccessDestination: + """Route succeeded item records to an S3 prefix.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return SuccessDestination( + bucket=bucket, + prefix=prefix or "", + include_input=include_input, + include_output=include_output, + expected_bucket_owner=expected_bucket_owner, + ) + + @staticmethod + def failures( + prefix_uri: str, + *, + include_input: bool = True, + include_error: bool = True, + expected_bucket_owner: str | None = None, + ) -> FailureDestination: + """Route permanently-failed item records to an S3 prefix.""" + parsed_uri = S3Uri.parse(prefix_uri) + bucket, prefix = parsed_uri.bucket, parsed_uri.path + return FailureDestination( + bucket=bucket, + prefix=prefix or "", + include_input=include_input, + include_error=include_error, + expected_bucket_owner=expected_bucket_owner, + ) + + +@dataclass(frozen=True) +class DistributedMapConfig: + """Configuration for map run operations.""" + + destination: DistributedMapDestinationConfig | None = None + completion_config: DistributedMapCompletionConfig | None = None + timeout: Duration | None = None + + def __post_init__(self) -> None: + if self.timeout is not None and not ( + 0 < self.timeout.to_seconds() <= 7776000 # noqa: PLR2004 + ): + msg = ( + "timeout must be positive and at most 90 days, got: " + f"{self.timeout.to_seconds()}s" + ) + raise ValidationError(msg) + + +@dataclass(frozen=True) +class DistributedMapResultConfig(DistributedMapConfig): + """Configuration that collects per-item results inline. + + Passing this instead of ``DistributedMapConfig`` makes ``ctx.distributed_map`` + return a ``DistributedMapResult``, which carries the per-item outcomes. + """ + + result_serdes: SerDes | None = None # None = DEFAULT_JSON_SERDES + + +# endregion map run configuration + + @dataclass(frozen=True) class ChildConfig(Generic[T]): """Configuration options for child context operations. diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/context.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/context.py index 16f5cced..e7de3162 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/context.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/context.py @@ -12,6 +12,7 @@ NoReturn, ParamSpec, TypeVar, + overload, ) from aws_durable_execution_sdk_python.config import ( @@ -20,6 +21,10 @@ Duration, InvokeConfig, MapConfig, + DistributedMapConfig, + DistributedMapProcessor, + DistributedMapResultConfig, + DistributedMapSource, ParallelBranch, ParallelConfig, StepConfig, @@ -28,6 +33,10 @@ from aws_durable_execution_sdk_python.concurrency.models import ( envelope_summary_generator, ) +from aws_durable_execution_sdk_python.dmap.models import ( + DistributedMapResult, + DistributedMapSummary, +) from aws_durable_execution_sdk_python.exceptions import ( CallbackError, CallbackExternalError, @@ -54,6 +63,9 @@ from aws_durable_execution_sdk_python.operation.child import child_handler from aws_durable_execution_sdk_python.operation.invoke import InvokeOperationExecutor from aws_durable_execution_sdk_python.operation.map import map_handler +from aws_durable_execution_sdk_python.operation.dmap import ( + DistributedMapOperationExecutor, +) from aws_durable_execution_sdk_python.operation.parallel import parallel_handler from aws_durable_execution_sdk_python.operation.step import StepOperationExecutor from aws_durable_execution_sdk_python.operation.wait import WaitOperationExecutor @@ -102,6 +114,8 @@ PASS_THROUGH_SERDES: SerDes[Any] = PassThroughSerDes() +_MAX_CONCURRENCY_LIMIT = 10000 + @dataclass(frozen=True) class ExecutionContext: @@ -683,6 +697,74 @@ def invoke( ) return executor.process() + @overload + def distributed_map( + self, + source: DistributedMapSource | Sequence[Any], + processor: DistributedMapProcessor, + max_concurrency: int, + name: str | None = None, + *, + config: DistributedMapResultConfig, + ) -> DistributedMapResult: ... + + @overload + def distributed_map( + self, + source: DistributedMapSource | Sequence[Any], + processor: DistributedMapProcessor, + max_concurrency: int, + name: str | None = None, + *, + config: DistributedMapConfig | None = None, + ) -> DistributedMapSummary: ... + + def distributed_map( + self, + source: DistributedMapSource | Sequence[Any], + processor: DistributedMapProcessor, + max_concurrency: int, + name: str | None = None, + *, + config: DistributedMapConfig | None = None, + ) -> DistributedMapSummary: + """Start a distributed map run and resolve with its summary. + + Args: + source: The items to process (a typed source or a plain-list shorthand) + processor: The processor configuration built via a DistributedMapProcessor factory + max_concurrency: Maximum concurrent processor invocations + name: Optional name for the operation + config: Optional run-level configuration + + Returns: + The map run's summary, or a DistributedMapResult when a + DistributedMapResultConfig is passed + """ + if not isinstance(source, (DistributedMapSource, list, tuple)): + msg = "source must be a DistributedMapSource or a list/tuple of items" + raise ValidationError(msg) + if not 1 <= max_concurrency <= _MAX_CONCURRENCY_LIMIT: + msg = ( + f"max_concurrency must be between 1 and {_MAX_CONCURRENCY_LIMIT}, " + f"got: {max_concurrency}" + ) + raise ValidationError(msg) + if config is None: + config = DistributedMapConfig() + with self._operation_replay_aware( + OperationSubType.DISTRIBUTED_MAP, name + ) as operation_identifier: + executor: DistributedMapOperationExecutor = DistributedMapOperationExecutor( + source=source, + processor=processor, + max_concurrency=max_concurrency, + state=self.state, + operation_identifier=operation_identifier, + config=config, + ) + return executor.process() + def map( self, inputs: Sequence[U], diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/__init__.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/handlers.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/handlers.py new file mode 100644 index 00000000..952744f7 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/handlers.py @@ -0,0 +1,453 @@ +"""Authoring helpers for distributed map processor and reader Lambda functions. + +Wrappers that own the request/response format so customers can write a plain +function to process items/batches or read pages. Imported explicitly from this +module, not the main package. +""" + +from __future__ import annotations + +import functools +from collections.abc import MutableMapping +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass, field +from typing import Any, Callable, Literal + +from aws_durable_execution_sdk_python.concurrency.models import BatchItemStatus +from aws_durable_execution_sdk_python.config import CompletionConfig, MapConfig +from aws_durable_execution_sdk_python.exceptions import ValidationError +from aws_durable_execution_sdk_python.execution import durable_execution +from aws_durable_execution_sdk_python.serdes import ( + DEFAULT_JSON_SERDES, + SerDes, + SerDesContext, +) + +_READER_STATE_LIMIT = 32 * 1024 + + +@dataclass(frozen=True) +class ReaderPage: + """A page returned by a reader function: its items and the next state.""" + + items: list[Any] = field(default_factory=list) + next_state: Any | None = None + + +@dataclass(frozen=True) +class ProcessorRecord: + """One item of the batch handed to a processor function.""" + + item_id: str + body: Any = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ProcessorRecord: + return cls(item_id=data.get("itemId", ""), body=data.get("body")) + + def to_dict(self) -> MutableMapping[str, Any]: + return {"itemId": self.item_id, "body": self.body} + + +@dataclass(frozen=True) +class ProcessorEvent: + """The event a processor function is invoked with, one batch of items.""" + + records: tuple[ProcessorRecord, ...] = () + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ProcessorEvent: + records = data.get("records") + if not isinstance(records, list): + msg = "expected a distributed map processor envelope with a 'records' list" + raise ValidationError(msg) + return cls(records=tuple(ProcessorRecord.from_dict(r) for r in records)) + + def to_dict(self) -> MutableMapping[str, Any]: + return {"records": [record.to_dict() for record in self.records]} + + +@dataclass(frozen=True) +class ItemResult: + """One item's output, reported by an item_results processor.""" + + item_identifier: str + output: Any = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ItemResult: + return cls( + item_identifier=data.get("itemIdentifier", ""), output=data.get("output") + ) + + def to_dict(self) -> MutableMapping[str, Any]: + return {"itemIdentifier": self.item_identifier, "output": self.output} + + +@dataclass(frozen=True) +class ItemFailure: + """One item's failure, reported by an item_failures or item_results processor.""" + + item_identifier: str + error_type: str = "" + error_message: str = "" + + @classmethod + def from_exception(cls, item_identifier: str, exc: BaseException) -> ItemFailure: + return cls( + item_identifier=item_identifier, + error_type=type(exc).__name__, + error_message=str(exc), + ) + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ItemFailure: + error = data.get("error") or {} + return cls( + item_identifier=data.get("itemIdentifier", ""), + error_type=error.get("errorType", ""), + error_message=error.get("errorMessage", ""), + ) + + def to_dict(self) -> MutableMapping[str, Any]: + return { + "itemIdentifier": self.item_identifier, + "error": { + "errorType": self.error_type, + "errorMessage": self.error_message, + }, + } + + +@dataclass(frozen=True) +class ItemHandlerResponse: + """What an item handler returns. A None ``results`` reports failures only.""" + + failures: tuple[ItemFailure, ...] = () + results: tuple[ItemResult, ...] | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ItemHandlerResponse: + results = data.get("batchItemResults") + return cls( + failures=tuple( + ItemFailure.from_dict(f) for f in data.get("batchItemFailures", []) + ), + results=tuple(ItemResult.from_dict(r) for r in results) + if results is not None + else None, + ) + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = { + "batchItemFailures": [f.to_dict() for f in self.failures] + } + if self.results is not None: + result["batchItemResults"] = [r.to_dict() for r in self.results] + return result + + +@dataclass(frozen=True) +class ReaderEvent: + """The event a reader function is invoked with.""" + + max_items: int + state: str | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ReaderEvent: + max_items = data.get("maxItems") + if not isinstance(max_items, int): + msg = ( + "expected a distributed map reader envelope with an integer 'maxItems'" + ) + raise ValidationError(msg) + return cls(max_items=max_items, state=data.get("state")) + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"maxItems": self.max_items} + if self.state is not None: + result["state"] = self.state + return result + + +@dataclass(frozen=True) +class ReaderResponse: + """What a reader function returns. A None ``next_state`` exhausts the source.""" + + items: tuple[Any, ...] = () + next_state: str | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> ReaderResponse: + return cls(items=tuple(data.get("items", ())), next_state=data.get("nextState")) + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"items": list(self.items)} + if self.next_state is not None: + result["nextState"] = self.next_state + return result + + +def _to_item(serdes: SerDes[Any], body: str, ctx: SerDesContext) -> Any: + """Recover a typed item from a record body.""" + return serdes.deserialize(body, ctx) + + +def _to_output(serdes: SerDes[Any], value: Any, ctx: SerDesContext) -> str: + """Serialize an item's output for the Output field.""" + return serdes.serialize(value, ctx) + + +def _validate_report(report: str) -> None: + if report not in ("results", "failures"): + msg = f"report must be 'results' or 'failures', got: {report!r}" + raise ValidationError(msg) + + +def _validate_concurrency(concurrency: int) -> None: + if concurrency < 1: + msg = f"concurrency must be at least 1, got: {concurrency}" + raise ValidationError(msg) + + +def distributed_map_item_handler( + func: Callable[[Any], Any] | None = None, + *, + item_serdes: SerDes[Any] | None = None, + result_serdes: SerDes[Any] | None = None, + concurrency: int = 1, + report: Literal["results", "failures"] = "results", +) -> Callable[..., Any]: + """Wrap a Lambda to be used as an item_results processor. + + Pass ``report="failures"`` for an item_failures processor. ``func`` + takes one item and returns its output or raises. The decorated name becomes + the Lambda handler and takes ``(event, context)``, so name it ``handler``. + + Items in a batch are processed one at a time. Raise ``concurrency`` to run + that many items at once, which requires ``func`` to be safe to call from + several threads. + """ + if func is None: + return functools.partial( + distributed_map_item_handler, + item_serdes=item_serdes, + result_serdes=result_serdes, + concurrency=concurrency, + report=report, + ) + _validate_report(report) + _validate_concurrency(concurrency) + serdes = item_serdes or DEFAULT_JSON_SERDES + out_serdes = result_serdes or DEFAULT_JSON_SERDES + + def handler(event: dict[str, Any], _context: Any = None) -> dict[str, Any]: + ctx = SerDesContext() + records = ProcessorEvent.from_dict(event).records + + def run(record: ProcessorRecord) -> Any: + return func(_to_item(serdes, record.body, ctx)) + + outputs: list[Any] = [None] * len(records) + errors: list[BaseException | None] = [None] * len(records) + with ThreadPoolExecutor(max_workers=concurrency) as pool: + futures = {pool.submit(run, r): i for i, r in enumerate(records)} + for future, i in futures.items(): + try: + outputs[i] = future.result() + except Exception as exc: # noqa: BLE001 + errors[i] = exc + + results: list[ItemResult] = [] + failures: list[ItemFailure] = [] + for i, record in enumerate(records): + err = errors[i] + if err is not None: + failures.append(ItemFailure.from_exception(record.item_id, err)) + elif report == "results": + results.append( + ItemResult( + item_identifier=record.item_id, + output=_to_output(out_serdes, outputs[i], ctx), + ) + ) + + response = ItemHandlerResponse( + failures=tuple(failures), + results=tuple(results) if report == "results" else None, + ) + return dict(response.to_dict()) + + return handler + + +def distributed_map_batch_handler( + func: Callable[[list[Any]], Any] | None = None, + *, + item_serdes: SerDes[Any] | None = None, +) -> Callable[..., Any]: + """Wrap a Lambda to be used as a batch processor. + + ``func`` takes the whole batch of items. Returning succeeds every item; + raising fails every item. The decorated name becomes the Lambda handler and + takes ``(event, context)``, so name it ``handler``. + """ + if func is None: + return functools.partial(distributed_map_batch_handler, item_serdes=item_serdes) + serdes = item_serdes or DEFAULT_JSON_SERDES + + def handler(event: dict[str, Any], _context: Any = None) -> Any: + ctx = SerDesContext() + return func( + [ + _to_item(serdes, record.body, ctx) + for record in ProcessorEvent.from_dict(event).records + ] + ) + + return handler + + +def distributed_map_reader( + func: Callable[[Any], ReaderPage] | None = None, + *, + state_serdes: SerDes[Any] | None = None, +) -> Callable[..., Any]: + """Wrap a Lambda to be used as a reader source. + + ``func`` takes the current state and returns a ReaderPage. A ``next_state`` + of ``None`` signals the source is exhausted. The decorated name becomes the + Lambda handler and takes ``(event, context)``, so name it ``handler``. + """ + if func is None: + return functools.partial(distributed_map_reader, state_serdes=state_serdes) + serdes = state_serdes or DEFAULT_JSON_SERDES + + def handler(event: dict[str, Any], _context: Any = None) -> dict[str, Any]: + ctx = SerDesContext() + reader_event = ReaderEvent.from_dict(event) + state = ( + serdes.deserialize(reader_event.state, ctx) + if reader_event.state is not None + else None + ) + + page = func(state) + if len(page.items) > reader_event.max_items: + msg = ( + f"reader returned {len(page.items)} items, exceeding maxItems " + f"{reader_event.max_items}" + ) + raise ValidationError(msg) + + next_state: str | None = None + if page.next_state is not None: + next_state = serdes.serialize(page.next_state, ctx) + if len(next_state.encode("utf-8")) > _READER_STATE_LIMIT: + msg = f"reader next_state exceeds the {_READER_STATE_LIMIT // 1024} KB limit" + raise ValidationError(msg) + + return dict( + ReaderResponse(items=tuple(page.items), next_state=next_state).to_dict() + ) + + return handler + + +def durable_distributed_map_item_handler( + func: Callable[..., Any] | None = None, + *, + item_serdes: SerDes[Any] | None = None, + result_serdes: SerDes[Any] | None = None, + report: Literal["results", "failures"] = "results", +) -> Callable[..., Any]: + """Durable variant of the item handler. ``func`` receives (context, item). + + The Lambda must be deployed as a durable function. The decorated name becomes + the Lambda handler and takes ``(event, context)``, so name it ``handler``. + """ + if func is None: + return functools.partial( + durable_distributed_map_item_handler, + item_serdes=item_serdes, + result_serdes=result_serdes, + report=report, + ) + _validate_report(report) + serdes = item_serdes or DEFAULT_JSON_SERDES + out_serdes = result_serdes or DEFAULT_JSON_SERDES + + @durable_execution + def handler(event: dict[str, Any], context: Any) -> dict[str, Any]: + ctx = SerDesContext() + records = ProcessorEvent.from_dict(event).records + + def per_item(inner_ctx: Any, body: Any, _index: int, _inputs: Any) -> Any: + return func(inner_ctx, _to_item(serdes, body, ctx)) + + batch = context.map( + [record.body for record in records], + per_item, + config=MapConfig(completion_config=CompletionConfig.all_completed()), + ) + + results: list[ItemResult] = [] + failures: list[ItemFailure] = [] + for bi in batch.all: + item_id = records[bi.index].item_id + if bi.status is BatchItemStatus.SUCCEEDED: + if report == "results": + results.append( + ItemResult( + item_identifier=item_id, + output=_to_output(out_serdes, bi.result, ctx), + ) + ) + else: + err = bi.error + failures.append( + ItemFailure( + item_identifier=item_id, + error_type=(err.type or "") if err else "", + error_message=(err.message or "") if err else "", + ) + ) + + response = ItemHandlerResponse( + failures=tuple(failures), + results=tuple(results) if report == "results" else None, + ) + return dict(response.to_dict()) + + return handler + + +def durable_distributed_map_batch_handler( + func: Callable[..., Any] | None = None, + *, + item_serdes: SerDes[Any] | None = None, +) -> Callable[..., Any]: + """Durable variant of the batch handler. ``func`` receives (context, items). + + The Lambda must be deployed as a durable function. The decorated name becomes + the Lambda handler and takes ``(event, context)``, so name it ``handler``. + """ + if func is None: + return functools.partial( + durable_distributed_map_batch_handler, item_serdes=item_serdes + ) + serdes = item_serdes or DEFAULT_JSON_SERDES + + @durable_execution + def handler(event: dict[str, Any], context: Any) -> Any: + ctx = SerDesContext() + return func( + context, + [ + _to_item(serdes, record.body, ctx) + for record in ProcessorEvent.from_dict(event).records + ], + ) + + return handler diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/models.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/models.py new file mode 100644 index 00000000..7d9ad80c --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/dmap/models.py @@ -0,0 +1,127 @@ +"""Result models returned by ``ctx.distributed_map``.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any + +from aws_durable_execution_sdk_python.exceptions import ( + DistributedMapError, +) +from aws_durable_execution_sdk_python.config import ( + DistributedMapCompletionReason, + DistributedMapItemStatus, + DistributedMapStatus, +) + + +@dataclass(frozen=True) +class DistributedMapSummary: + """Outcome of a map run without per-item results. + + Resolved by ``ctx.distributed_map`` for every terminal state. A non-``SUCCEEDED`` + run resolves rather than raising. Use :meth:`throw_if_error` to opt into raising. + """ + + status: DistributedMapStatus + completion_reason: DistributedMapCompletionReason + success_count: int + failure_count: int + unprocessed_count: int + distributed_map_run_arn: str | None = None + completion_details: str | None = None + total_count: int | None = None + + @property + def has_failure(self) -> bool: + """``True`` when any item permanently failed.""" + return self.failure_count > 0 + + def throw_if_error(self) -> None: + """Raise :class:`DistributedMapError` on any non-success outcome.""" + if self.status is not DistributedMapStatus.SUCCEEDED: + detail = f", {self.completion_details}" if self.completion_details else "" + msg = ( + f"Map run ended {self.status.value} " + f"(reason: {self.completion_reason.value}{detail})" + ) + raise DistributedMapError(msg) + if self.failure_count > 0: + msg = ( + f"Map run succeeded but {self.failure_count} item(s) permanently failed" + ) + raise DistributedMapError(msg) + + +@dataclass(frozen=True) +class DistributedMapItemError: + """Error for a single failed map run item.""" + + error_type: str + error_message: str + + +@dataclass(frozen=True) +class DistributedMapResultItem: + """Outcome of a single map run item.""" + + item_id: str + status: DistributedMapItemStatus + output: Any | None = None + error: DistributedMapItemError | None = None + + +@dataclass(frozen=True) +class DistributedMapResult(DistributedMapSummary): + """Outcome of a map run with per-item results. + + Resolved by ``ctx.distributed_map`` when passed a + :class:`DistributedMapResultConfig`. Extends :class:`DistributedMapSummary` + with the retained per-item results. + """ + + all: list[DistributedMapResultItem] = field(default_factory=list) + + def succeeded(self) -> list[DistributedMapResultItem]: + """Return the items that succeeded.""" + return [ + item + for item in self.all + if item.status is DistributedMapItemStatus.SUCCEEDED + ] + + def failed(self) -> list[DistributedMapResultItem]: + """Return the items that permanently failed.""" + return [ + item for item in self.all if item.status is DistributedMapItemStatus.FAILED + ] + + def get_results(self) -> list[Any]: + """Return the outputs of the succeeded items.""" + return [ + item.output + for item in self.all + if item.status is DistributedMapItemStatus.SUCCEEDED + and item.output is not None + ] + + def get_errors(self) -> list[DistributedMapItemError]: + """Return the errors of the failed items.""" + return [ + item.error + for item in self.all + if item.status is DistributedMapItemStatus.FAILED and item.error is not None + ] + + def throw_if_error(self) -> None: + """Raise the first failed item's error, otherwise defer to the summary rule.""" + if self.status is DistributedMapStatus.SUCCEEDED and self.failure_count > 0: + failed = self.failed() + if failed: + first = failed[0] + if first.error is not None: + msg = f"{first.error.error_type}: {first.error.error_message}" + raise DistributedMapError(msg) + msg = f"item {first.item_id} failed" + raise DistributedMapError(msg) + super().throw_if_error() diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/exceptions.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/exceptions.py index ea3f163f..b2123d7a 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/exceptions.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/exceptions.py @@ -330,6 +330,10 @@ class InvokeError(DurableOperationError): """Raised when a durable invoke operation fails.""" +class DistributedMapError(DurableOperationError): + """Raised when a durable map run operation fails.""" + + class ChildContextError(DurableOperationError): """Raised when a child context (run_in_child_context, map, parallel) fails.""" diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py index eb5ee78b..a51568d1 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/lambda_service.py @@ -7,15 +7,31 @@ from collections.abc import MutableMapping from dataclasses import dataclass, field from enum import Enum -from typing import TYPE_CHECKING, Any, NoReturn, Protocol, TypeAlias, cast +from typing import ( + TYPE_CHECKING, + Any, + NoReturn, + Protocol, + Self, + TypeAlias, + TypeVar, + cast, +) import boto3 from botocore.config import Config from aws_durable_execution_sdk_python.__about__ import __version__ +from aws_durable_execution_sdk_python.config import ( + DistributedMapCompletionReason, + DistributedMapCsvDelimiter, + DistributedMapItemStatus, + DistributedMapSourceFormat, +) from aws_durable_execution_sdk_python.exceptions import ( CheckpointError, DurableOperationError, + ExecutionError, GetExecutionStateError, SerDesError, ) @@ -74,6 +90,7 @@ class OperationType(Enum): WAIT = "WAIT" CALLBACK = "CALLBACK" CHAINED_INVOKE = "CHAINED_INVOKE" + DISTRIBUTED_MAP = "DISTRIBUTED_MAP" @classmethod def from_sub_type(cls, sub_type: OperationSubType) -> OperationType: @@ -86,6 +103,8 @@ def from_sub_type(cls, sub_type: OperationSubType) -> OperationType: return OperationType.CHAINED_INVOKE case OperationSubType.CALLBACK: return OperationType.CALLBACK + case OperationSubType.DISTRIBUTED_MAP: + return OperationType.DISTRIBUTED_MAP case ( OperationSubType.WAIT_FOR_CALLBACK | OperationSubType.RUN_IN_CHILD_CONTEXT @@ -104,6 +123,42 @@ class CallbackTimeoutType(Enum): HEARTBEAT = "Callback.Heartbeat" +class DistributedMapSourceType(Enum): + INLINE = "INLINE" + S3 = "S3" + READER_FUNCTION = "READER_FUNCTION" + + +class DistributedMapS3SourceTransform(Enum): + NONE = "NONE" + LOAD_AND_FLATTEN = "LOAD_AND_FLATTEN" + + +class DistributedMapCsvHeaderLocation(Enum): + FIRST_ROW = "FIRST_ROW" + GIVEN = "GIVEN" + + +class DistributedMapFunctionResponseType(Enum): + REPORT_BATCH_ITEM_FAILURES = "REPORT_BATCH_ITEM_FAILURES" + REPORT_BATCH_ITEM_RESULTS = "REPORT_BATCH_ITEM_RESULTS" + + +class DistributedMapDestinationType(Enum): + S3 = "S3" + + +class DistributedMapDestinationInclude(Enum): + INPUT = "INPUT" + OUTPUT = "OUTPUT" + ERROR = "ERROR" + + +class DistributedMapResultCollectionMode(Enum): + NONE = "NONE" + INLINE = "INLINE" + + class OperationSubType(Enum): STEP = "Step" WAIT = "Wait" @@ -116,6 +171,7 @@ class OperationSubType(Enum): WAIT_FOR_CALLBACK = "WaitForCallback" WAIT_FOR_CONDITION = "WaitForCondition" CHAINED_INVOKE = "ChainedInvoke" + DISTRIBUTED_MAP = "DistributedMap" class InvocationStatus(Enum): @@ -383,6 +439,101 @@ def from_dict(cls, data: MutableMapping[str, Any]) -> ChainedInvokeDetails: ) +_BackendEnumT = TypeVar("_BackendEnumT", bound=Enum) + + +def _parse_enum( + enum_cls: type[_BackendEnumT], + value: str, + field_name: str, + *, + unknown_fallback: _BackendEnumT | None = None, +) -> _BackendEnumT: + """Convert a backend enum string, returning unknown_fallback for an unrecognized value or raising ExecutionError when no fallback is given.""" + try: + return enum_cls(value) + except ValueError as e: + if unknown_fallback is not None: + return unknown_fallback + msg = f"Unknown distributed map {field_name} from the backend: {value!r}" + raise ExecutionError(msg) from e + + +@dataclass(frozen=True) +class DistributedMapResultItem: + """Represent a single map run item's outcome.""" + + item_id: str + status: DistributedMapItemStatus + output: str | None = None + error: ErrorObject | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapResultItem: + error_raw = data.get("Error") + return cls( + item_id=data.get("ItemId", ""), + status=_parse_enum( + DistributedMapItemStatus, data.get("Status", ""), "item status" + ), + output=data.get("Output"), + error=ErrorObject.from_dict(error_raw) if error_raw else None, + ) + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = { + "ItemId": self.item_id, + "Status": self.status.value, + } + if self.output is not None: + result["Output"] = self.output + if self.error is not None: + result["Error"] = self.error.to_dict() + return result + + +@dataclass(frozen=True) +class DistributedMapDetails: + completion_reason: DistributedMapCompletionReason | None = None + distributed_map_run_arn: str | None = None + completion_details: str | None = None + total_count: int | None = None + success_count: int = 0 + failure_count: int = 0 + unprocessed_count: int = 0 + results: tuple[DistributedMapResultItem, ...] | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapDetails: + results_raw = data.get("Results") + reason_raw = data.get("CompletionReason") + # CompletionReason is absent while a run is in flight, so parse it only + # when present rather than requiring it. + return cls( + completion_reason=( + _parse_enum( + DistributedMapCompletionReason, + reason_raw, + "completion reason", + unknown_fallback=DistributedMapCompletionReason.UNKNOWN_TO_SDK_VERSION, + ) + if reason_raw is not None + else None + ), + distributed_map_run_arn=data.get("DistributedMapRunArn"), + completion_details=data.get("CompletionDetails"), + total_count=data.get("TotalCount"), + success_count=data.get("SuccessCount", 0), + failure_count=data.get("FailureCount", 0), + unprocessed_count=data.get("UnprocessedCount", 0), + results=tuple( + DistributedMapResultItem.from_dict(item) for item in results_raw + ) + if results_raw is not None + else None, + ) + + @dataclass(frozen=True) class StepOptions: next_attempt_delay_seconds: int = 0 @@ -478,6 +629,419 @@ def to_dict(self) -> MutableMapping[str, Any]: return result +@dataclass(frozen=True) +class DistributedMapCsvFormatOptions: + """Represent CSV format options for an S3 source.""" + + header_location: DistributedMapCsvHeaderLocation + headers: tuple[str, ...] | None = None + delimiter: DistributedMapCsvDelimiter | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = { + "HeaderLocation": self.header_location.value + } + if self.headers is not None: + result["Headers"] = list(self.headers) + if self.delimiter is not None: + result["Delimiter"] = self.delimiter.value + return result + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapCsvFormatOptions: + headers = data.get("Headers") + delimiter = data.get("Delimiter") + return cls( + header_location=DistributedMapCsvHeaderLocation( + data.get("HeaderLocation", "FIRST_ROW") + ), + headers=tuple(headers) if headers is not None else None, + delimiter=DistributedMapCsvDelimiter(delimiter) + if delimiter is not None + else None, + ) + + +@dataclass(frozen=True) +class DistributedMapS3SourceConfig: + """Represent an S3 distributed map source config.""" + + bucket: str + key: str | None = None + key_prefix: str | None = None + transform: DistributedMapS3SourceTransform | None = None + expected_bucket_owner: str | None = None + fmt: DistributedMapSourceFormat | None = None + csv_format_options: DistributedMapCsvFormatOptions | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"Bucket": self.bucket} + if self.key is not None: + result["Key"] = self.key + if self.key_prefix is not None: + result["KeyPrefix"] = self.key_prefix + if self.transform is not None: + result["Transform"] = self.transform.value + if self.expected_bucket_owner is not None: + result["ExpectedBucketOwner"] = self.expected_bucket_owner + if self.fmt is not None: + result["Format"] = self.fmt.value + if self.csv_format_options is not None: + result["CsvFormatOptions"] = self.csv_format_options.to_dict() + return result + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapS3SourceConfig: + transform = data.get("Transform") + fmt = data.get("Format") + csv_raw = data.get("CsvFormatOptions") + return cls( + bucket=data.get("Bucket", ""), + key=data.get("Key"), + key_prefix=data.get("KeyPrefix"), + transform=DistributedMapS3SourceTransform(transform) + if transform is not None + else None, + expected_bucket_owner=data.get("ExpectedBucketOwner"), + fmt=DistributedMapSourceFormat(fmt) if fmt is not None else None, + csv_format_options=DistributedMapCsvFormatOptions.from_dict(csv_raw) + if csv_raw is not None + else None, + ) + + +@dataclass(frozen=True) +class DistributedMapReaderFunctionSourceConfig: + """Represent a reader-function distributed map source config.""" + + function_name: str + initial_state: str | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"FunctionName": self.function_name} + if self.initial_state is not None: + result["InitialState"] = self.initial_state + return result + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapReaderFunctionSourceConfig: + return cls( + function_name=data.get("FunctionName", ""), + initial_state=data.get("InitialState"), + ) + + +@dataclass(frozen=True) +class DistributedMapInlineSourceConfig: + """Represent the items an inline map run source reads from.""" + + items: tuple[str, ...] = () + + def to_dict(self) -> MutableMapping[str, Any]: + return {"Items": list(self.items)} + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapInlineSourceConfig: + return cls(items=tuple(data.get("Items", ()))) + + +@dataclass(frozen=True) +class DistributedMapSourceConfig: + """Represent a map run source.""" + + source_type: DistributedMapSourceType + max_items: int | None = None + inline_source_config: DistributedMapInlineSourceConfig | None = None + s3_config: DistributedMapS3SourceConfig | None = None + reader_config: DistributedMapReaderFunctionSourceConfig | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"Type": self.source_type.value} + if self.source_type is DistributedMapSourceType.INLINE: + inline = self.inline_source_config or DistributedMapInlineSourceConfig() + result["InlineSourceConfig"] = inline.to_dict() + elif ( + self.source_type is DistributedMapSourceType.S3 + and self.s3_config is not None + ): + result["S3SourceConfig"] = self.s3_config.to_dict() + elif ( + self.source_type is DistributedMapSourceType.READER_FUNCTION + and self.reader_config is not None + ): + result["ReaderFunctionSourceConfig"] = self.reader_config.to_dict() + if self.max_items is not None: + result["MaxItemsToRead"] = self.max_items + return result + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapSourceConfig: + source_type = DistributedMapSourceType(data.get("Type", "INLINE")) + inline_cfg = data.get("InlineSourceConfig") or {} + s3_raw = data.get("S3SourceConfig") + reader_raw = data.get("ReaderFunctionSourceConfig") + return cls( + source_type=source_type, + max_items=data.get("MaxItemsToRead"), + inline_source_config=DistributedMapInlineSourceConfig.from_dict(inline_cfg) + if source_type is DistributedMapSourceType.INLINE + else None, + s3_config=DistributedMapS3SourceConfig.from_dict(s3_raw) + if s3_raw is not None + else None, + reader_config=DistributedMapReaderFunctionSourceConfig.from_dict(reader_raw) + if reader_raw is not None + else None, + ) + + +@dataclass(frozen=True) +class DistributedMapProcessorConfig: + """Represent a map run processor.""" + + function_name: str + function_response_types: tuple[DistributedMapFunctionResponseType, ...] | None = ( + None + ) + batch_size: int | None = None + max_retry_attempts: int | None = None + max_retry_duration_seconds: int | None = None + durable_execution_name_prefix: str | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {"FunctionName": self.function_name} + if self.function_response_types: + result["FunctionResponseTypes"] = [ + t.value for t in self.function_response_types + ] + if self.batch_size is not None: + result["BatchSize"] = self.batch_size + if self.max_retry_attempts is not None: + result["MaxRetryAttempts"] = self.max_retry_attempts + if self.max_retry_duration_seconds is not None: + result["MaxRetryDurationSeconds"] = self.max_retry_duration_seconds + if self.durable_execution_name_prefix is not None: + result["DurableExecutionNamePrefix"] = self.durable_execution_name_prefix + return result + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapProcessorConfig: + response_types = data.get("FunctionResponseTypes") + return cls( + function_name=data.get("FunctionName", ""), + function_response_types=tuple( + DistributedMapFunctionResponseType(t) for t in response_types + ) + if response_types + else None, + batch_size=data.get("BatchSize"), + max_retry_attempts=data.get("MaxRetryAttempts"), + max_retry_duration_seconds=data.get("MaxRetryDurationSeconds"), + durable_execution_name_prefix=data.get("DurableExecutionNamePrefix"), + ) + + +@dataclass(frozen=True) +class DistributedMapCompletionConfig: + """Represent a map run completion (failure-tolerance) config.""" + + tolerated_failure_count: int | None = None + tolerated_failure_percentage: float | None = None + minimum_sample_size: int | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {} + if self.tolerated_failure_count is not None: + result["ToleratedFailureCount"] = self.tolerated_failure_count + if self.tolerated_failure_percentage is not None: + result["ToleratedFailurePercentage"] = self.tolerated_failure_percentage + if self.minimum_sample_size is not None: + result["MinimumSampleSize"] = self.minimum_sample_size + return result + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapCompletionConfig: + return cls( + tolerated_failure_count=data.get("ToleratedFailureCount"), + tolerated_failure_percentage=data.get("ToleratedFailurePercentage"), + minimum_sample_size=data.get("MinimumSampleSize"), + ) + + +@dataclass(frozen=True) +class DistributedMapS3DestinationConfig: + """Represent an S3 distributed map destination config.""" + + bucket: str + key_prefix: str + expected_bucket_owner: str | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = { + "Bucket": self.bucket, + "KeyPrefix": self.key_prefix, + } + if self.expected_bucket_owner is not None: + result["ExpectedBucketOwner"] = self.expected_bucket_owner + return result + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapS3DestinationConfig: + return cls( + bucket=data.get("Bucket", ""), + key_prefix=data.get("KeyPrefix", ""), + expected_bucket_owner=data.get("ExpectedBucketOwner"), + ) + + +@dataclass(frozen=True) +class _DistributedMapDestinationEntry: + """Hold the members the OnSuccess and OnFailure shapes have in common.""" + + type: DistributedMapDestinationType + include: tuple[DistributedMapDestinationInclude, ...] + s3_destination_config: DistributedMapS3DestinationConfig + + def to_dict(self) -> MutableMapping[str, Any]: + return { + "Type": self.type.value, + "Include": [i.value for i in self.include], + "S3DestinationConfig": self.s3_destination_config.to_dict(), + } + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> Self: + s3_raw = data.get("S3DestinationConfig") or {} + return cls( + type=DistributedMapDestinationType(data.get("Type", "S3")), + include=tuple( + DistributedMapDestinationInclude(i) for i in data.get("Include", []) + ), + s3_destination_config=DistributedMapS3DestinationConfig.from_dict(s3_raw), + ) + + +@dataclass(frozen=True) +class DistributedMapOnSuccessConfig(_DistributedMapDestinationEntry): + """Represent where a map run writes the outcome of its succeeded items.""" + + +@dataclass(frozen=True) +class DistributedMapOnFailureConfig(_DistributedMapDestinationEntry): + """Represent where a map run writes the outcome of its failed items.""" + + +@dataclass(frozen=True) +class DistributedMapDestinationConfig: + """Represent a map run destination config.""" + + on_success: DistributedMapOnSuccessConfig | None = None + on_failure: DistributedMapOnFailureConfig | None = None + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = {} + if self.on_success is not None: + result["OnSuccess"] = self.on_success.to_dict() + if self.on_failure is not None: + result["OnFailure"] = self.on_failure.to_dict() + return result + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapDestinationConfig: + success_raw = data.get("OnSuccess") + failure_raw = data.get("OnFailure") + return cls( + on_success=DistributedMapOnSuccessConfig.from_dict(success_raw) + if success_raw is not None + else None, + on_failure=DistributedMapOnFailureConfig.from_dict(failure_raw) + if failure_raw is not None + else None, + ) + + +@dataclass(frozen=True) +class DistributedMapResultCollectionConfig: + """Represent the map run result-collection setting.""" + + mode: DistributedMapResultCollectionMode + + def to_dict(self) -> MutableMapping[str, Any]: + return {"Mode": self.mode.value} + + @classmethod + def from_dict( + cls, data: MutableMapping[str, Any] + ) -> DistributedMapResultCollectionConfig: + return cls(mode=DistributedMapResultCollectionMode(data.get("Mode", "NONE"))) + + +@dataclass(frozen=True) +class DistributedMapOptions: + """Configuration options for starting a map run.""" + + max_concurrency: int + source: DistributedMapSourceConfig + processor: DistributedMapProcessorConfig + destination: DistributedMapDestinationConfig | None = None + completion_config: DistributedMapCompletionConfig | None = None + result_collection: DistributedMapResultCollectionConfig | None = None + timeout_seconds: int | None = None + + @classmethod + def from_dict(cls, data: MutableMapping[str, Any]) -> DistributedMapOptions: + source_raw = data.get("Source") or {} + processor_raw = data.get("Processor") or {} + destination_raw = data.get("Destination") + completion_raw = data.get("CompletionConfig") + result_collection_raw = data.get("ResultCollection") + return cls( + max_concurrency=data["MaxConcurrency"], + source=DistributedMapSourceConfig.from_dict(source_raw), + processor=DistributedMapProcessorConfig.from_dict(processor_raw), + destination=DistributedMapDestinationConfig.from_dict(destination_raw) + if destination_raw is not None + else None, + completion_config=DistributedMapCompletionConfig.from_dict(completion_raw) + if completion_raw is not None + else None, + result_collection=DistributedMapResultCollectionConfig.from_dict( + result_collection_raw + ) + if result_collection_raw is not None + else None, + timeout_seconds=data.get("TimeoutSeconds"), + ) + + def to_dict(self) -> MutableMapping[str, Any]: + result: MutableMapping[str, Any] = { + "MaxConcurrency": self.max_concurrency, + "Source": self.source.to_dict(), + "Processor": self.processor.to_dict(), + } + if self.destination is not None: + result["Destination"] = self.destination.to_dict() + if self.completion_config is not None: + result["CompletionConfig"] = self.completion_config.to_dict() + if self.result_collection is not None: + result["ResultCollection"] = self.result_collection.to_dict() + if self.timeout_seconds is not None: + result["TimeoutSeconds"] = self.timeout_seconds + return result + + @dataclass(frozen=True) class ContextOptions: replay_children: ReplayChildren = False @@ -510,6 +1074,7 @@ class OperationUpdate: wait_options: WaitOptions | None = None callback_options: CallbackOptions | None = None chained_invoke_options: ChainedInvokeOptions | None = None + distributed_map_options: DistributedMapOptions | None = None def to_dict(self) -> MutableMapping[str, Any]: result: MutableMapping[str, Any] = { @@ -538,6 +1103,8 @@ def to_dict(self) -> MutableMapping[str, Any]: result["CallbackOptions"] = self.callback_options.to_dict() if self.chained_invoke_options: result["ChainedInvokeOptions"] = self.chained_invoke_options.to_dict() + if self.distributed_map_options: + result["DistributedMapOptions"] = self.distributed_map_options.to_dict() return result @@ -566,6 +1133,12 @@ def from_dict(cls, data: MutableMapping[str, Any]) -> OperationUpdate: if invoke_data := data.get("ChainedInvokeOptions"): chained_invoke_options = ChainedInvokeOptions.from_dict(invoke_data) + distributed_map_options = None + if distributed_map_options_data := data.get("DistributedMapOptions"): + distributed_map_options = DistributedMapOptions.from_dict( + distributed_map_options_data + ) + return cls( operation_id=data["Id"], operation_type=OperationType(data["Type"]), @@ -580,6 +1153,7 @@ def from_dict(cls, data: MutableMapping[str, Any]) -> OperationUpdate: wait_options=wait_options, callback_options=callback_options, chained_invoke_options=chained_invoke_options, + distributed_map_options=distributed_map_options, ) @classmethod @@ -763,6 +1337,26 @@ def create_invoke_start( # endregion invoke + # region map run + @classmethod + def create_distributed_map_start( + cls, + identifier: OperationIdentifier, + distributed_map_options: DistributedMapOptions, + ) -> OperationUpdate: + """Create an instance of OperationUpdate for type: DISTRIBUTED_MAP, action: START.""" + return cls( + operation_id=identifier.operation_id, + parent_id=identifier.parent_id, + operation_type=OperationType.DISTRIBUTED_MAP, + sub_type=OperationSubType.DISTRIBUTED_MAP, + action=OperationAction.START, + name=identifier.name, + distributed_map_options=distributed_map_options, + ) + + # endregion map run + # region wait for condition @classmethod def create_wait_for_condition_start( @@ -886,6 +1480,7 @@ class Operation: wait_details: WaitDetails | None = None callback_details: CallbackDetails | None = None chained_invoke_details: ChainedInvokeDetails | None = None + distributed_map_details: DistributedMapDetails | None = None @classmethod def from_dict(cls, data: MutableMapping[str, Any]) -> Operation: @@ -930,6 +1525,12 @@ def from_dict(cls, data: MutableMapping[str, Any]) -> Operation: chained_invoke_details ) + distributed_map_details = None + if distributed_map_details_input := data.get("DistributedMapDetails"): + distributed_map_details = DistributedMapDetails.from_dict( + distributed_map_details_input + ) + return cls( operation_id=data["Id"], operation_type=operation_type, @@ -945,6 +1546,7 @@ def from_dict(cls, data: MutableMapping[str, Any]) -> Operation: wait_details=wait_details, callback_details=callback_details, chained_invoke_details=chained_invoke_details, + distributed_map_details=distributed_map_details, ) def to_dict(self) -> MutableMapping[str, Any]: @@ -1009,6 +1611,33 @@ def to_dict(self) -> MutableMapping[str, Any]: if self.chained_invoke_details.error: invoke_dict["Error"] = self.chained_invoke_details.error.to_dict() result["ChainedInvokeDetails"] = invoke_dict + if self.distributed_map_details: + distributed_map_details_dict: MutableMapping[str, Any] = { + "SuccessCount": self.distributed_map_details.success_count, + "FailureCount": self.distributed_map_details.failure_count, + "UnprocessedCount": self.distributed_map_details.unprocessed_count, + } + if self.distributed_map_details.completion_reason is not None: + distributed_map_details_dict["CompletionReason"] = ( + self.distributed_map_details.completion_reason.value + ) + if self.distributed_map_details.distributed_map_run_arn: + distributed_map_details_dict["DistributedMapRunArn"] = ( + self.distributed_map_details.distributed_map_run_arn + ) + if self.distributed_map_details.completion_details: + distributed_map_details_dict["CompletionDetails"] = ( + self.distributed_map_details.completion_details + ) + if self.distributed_map_details.total_count is not None: + distributed_map_details_dict["TotalCount"] = ( + self.distributed_map_details.total_count + ) + if self.distributed_map_details.results is not None: + distributed_map_details_dict["Results"] = [ + item.to_dict() for item in self.distributed_map_details.results + ] + result["DistributedMapDetails"] = distributed_map_details_dict return result def to_json_dict(self) -> MutableMapping[str, Any]: diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/dmap.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/dmap.py new file mode 100644 index 00000000..33ed11e1 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/operation/dmap.py @@ -0,0 +1,582 @@ +"""Implement the Durable map run operation.""" + +from __future__ import annotations + +import json +import logging +from typing import TYPE_CHECKING, Any + +from aws_durable_execution_sdk_python.config import ( + DistributedMapProcessor, + DistributedMapResultConfig, + DistributedMapSource, + DistributedMapSourceFormat, + DistributedMapStatus, + ProcessorResponseMode, +) +from aws_durable_execution_sdk_python.dmap.models import ( + DistributedMapItemError, + DistributedMapResult, + DistributedMapResultItem, + DistributedMapSummary, +) +from aws_durable_execution_sdk_python.exceptions import ( + ExecutionError, + ValidationError, +) +from aws_durable_execution_sdk_python.lambda_service import ( + DistributedMapCompletionConfig, + DistributedMapCsvFormatOptions, + DistributedMapCsvHeaderLocation, + DistributedMapDestinationConfig, + DistributedMapDestinationInclude, + DistributedMapDestinationType, + DistributedMapFunctionResponseType, + DistributedMapInlineSourceConfig, + DistributedMapOnFailureConfig, + DistributedMapOnSuccessConfig, + DistributedMapOptions, + DistributedMapProcessorConfig, + DistributedMapReaderFunctionSourceConfig, + DistributedMapResultCollectionConfig, + DistributedMapResultCollectionMode, + DistributedMapS3DestinationConfig, + DistributedMapS3SourceConfig, + DistributedMapS3SourceTransform, + DistributedMapSourceConfig, + DistributedMapSourceType, + OperationStatus, + OperationUpdate, +) +from aws_durable_execution_sdk_python.operation.base import ( + CheckResult, + OperationExecutor, +) +from aws_durable_execution_sdk_python.serdes import ( + DEFAULT_JSON_SERDES, + deserialize, + serialize, +) +from aws_durable_execution_sdk_python.suspend import suspend_with_optional_resume_delay + +if TYPE_CHECKING: + from collections.abc import Sequence + + from aws_durable_execution_sdk_python.config import ( + DistributedMapConfig, + ) + from aws_durable_execution_sdk_python.identifier import OperationIdentifier + from aws_durable_execution_sdk_python.lambda_service import ( + DistributedMapDetails, + Operation, + ) + from aws_durable_execution_sdk_python.state import ( + CheckpointedResult, + ExecutionState, + ) + +logger = logging.getLogger(__name__) + +# Size limits for the inline item list (1 MB) and the reader's saved state (32 KB). +_INLINE_SIZE_LIMIT = 1024 * 1024 +_READER_STATE_LIMIT = 32 * 1024 +_UNLIMITED_RETRY_WIRE = -1 + +_RESPONSE_TYPE_FOR_MODE = { + ProcessorResponseMode.ITEM_FAILURES: DistributedMapFunctionResponseType.REPORT_BATCH_ITEM_FAILURES, + ProcessorResponseMode.ITEM_RESULTS: DistributedMapFunctionResponseType.REPORT_BATCH_ITEM_RESULTS, +} + + +def _build_inline_items( + items: tuple[Any, ...], + serdes: Any, + operation_id: str, + durable_execution_arn: str, +) -> tuple[str, ...]: + """Serialize each inline item, enforcing the 1 MB cap on the whole list.""" + serialized_items: list[str] = [] + for item in items: + serialized_items.append( + serialize( + serdes=serdes, + value=item, + operation_id=operation_id, + durable_execution_arn=durable_execution_arn, + ) + ) + total = len(json.dumps(serialized_items, separators=(",", ":")).encode("utf-8")) + if total > _INLINE_SIZE_LIMIT: + msg = ( + f"inline source exceeds the {_INLINE_SIZE_LIMIT // 1024 // 1024} MB limit " + f"(serialized size: {total} bytes)" + ) + raise ValidationError(msg) + return tuple(serialized_items) + + +def _build_completion_config( + completion_config: Any, +) -> DistributedMapCompletionConfig | None: + """Translate the completion config, or None when it carries no threshold.""" + if ( + completion_config.tolerated_failure_count is None + and completion_config.tolerated_failure_percentage is None + and completion_config.minimum_sample_size is None + ): + return None + return DistributedMapCompletionConfig( + tolerated_failure_count=completion_config.tolerated_failure_count, + tolerated_failure_percentage=completion_config.tolerated_failure_percentage, + minimum_sample_size=completion_config.minimum_sample_size, + ) + + +def _build_processor_config( + processor: DistributedMapProcessor, +) -> DistributedMapProcessorConfig: + """Translate the processor, mapping the response mode and unlimited retries.""" + response_type = _RESPONSE_TYPE_FOR_MODE.get(processor.response_mode) + max_retry_attempts: int | None = None + max_retry_duration_seconds: int | None = None + attempts = processor.max_retry_attempts + if attempts == DistributedMapProcessor.UNLIMITED: + max_retry_attempts = _UNLIMITED_RETRY_WIRE + elif isinstance(attempts, int): + max_retry_attempts = attempts + if processor.max_retry_duration is not None: + max_retry_duration_seconds = processor.max_retry_duration.to_seconds() + return DistributedMapProcessorConfig( + function_name=processor.function_name, + function_response_types=(response_type,) if response_type else None, + batch_size=processor.batch_size, + max_retry_attempts=max_retry_attempts, + max_retry_duration_seconds=max_retry_duration_seconds, + durable_execution_name_prefix=processor.durable_execution_name_prefix, + ) + + +def _transform_for(s3: Any) -> DistributedMapS3SourceTransform | None: + """A prefix source flattens when a format is set, and lists keys when it is not.""" + if s3.prefix is None: + return None + if s3.fmt is None: + return DistributedMapS3SourceTransform.NONE + return DistributedMapS3SourceTransform.LOAD_AND_FLATTEN + + +def _build_s3_source_config(s3: Any) -> DistributedMapS3SourceConfig: + """Translate the S3 source, deriving the transform and CSV header location.""" + csv_format_options: DistributedMapCsvFormatOptions | None = None + if s3.fmt is DistributedMapSourceFormat.CSV: + csv_format_options = DistributedMapCsvFormatOptions( + header_location=( + DistributedMapCsvHeaderLocation.GIVEN + if s3.headers is not None + else DistributedMapCsvHeaderLocation.FIRST_ROW + ), + headers=s3.headers, + delimiter=s3.delimiter, + ) + return DistributedMapS3SourceConfig( + bucket=s3.bucket, + key=s3.key, + key_prefix=s3.prefix, + transform=_transform_for(s3), + expected_bucket_owner=s3.expected_bucket_owner, + fmt=s3.fmt, + csv_format_options=csv_format_options, + ) + + +def _destination_entry_fields( + destination: Any, *, second_include: DistributedMapDestinationInclude +) -> dict[str, Any]: + """Build the members the OnSuccess and OnFailure shapes share.""" + include: list[DistributedMapDestinationInclude] = [] + if destination.include_input: + include.append(DistributedMapDestinationInclude.INPUT) + if second_include is DistributedMapDestinationInclude.OUTPUT: + if destination.include_output: + include.append(second_include) + elif destination.include_error: + include.append(second_include) + return { + "type": DistributedMapDestinationType.S3, + "include": tuple(include), + "s3_destination_config": DistributedMapS3DestinationConfig( + bucket=destination.bucket, + key_prefix=destination.prefix, + expected_bucket_owner=destination.expected_bucket_owner, + ), + } + + +def _build_destination_config( + destination: Any, +) -> DistributedMapDestinationConfig | None: + """Translate the destination config, or None when neither side is set.""" + on_success = ( + DistributedMapOnSuccessConfig( + **_destination_entry_fields( + destination.on_success, + second_include=DistributedMapDestinationInclude.OUTPUT, + ) + ) + if destination.on_success is not None + else None + ) + on_failure = ( + DistributedMapOnFailureConfig( + **_destination_entry_fields( + destination.on_failure, + second_include=DistributedMapDestinationInclude.ERROR, + ) + ) + if destination.on_failure is not None + else None + ) + if on_success is None and on_failure is None: + return None + return DistributedMapDestinationConfig(on_success=on_success, on_failure=on_failure) + + +def _build_source_config( + source: DistributedMapSource | Sequence[Any], + operation_id: str, + durable_execution_arn: str, +) -> DistributedMapSourceConfig: + """Translate a source (typed or plain-list shorthand) into its source config.""" + if not isinstance(source, DistributedMapSource): + # A plain list is treated as an inline source with the default serializer. + serialized_items = _build_inline_items( + tuple(source), DEFAULT_JSON_SERDES, operation_id, durable_execution_arn + ) + return DistributedMapSourceConfig( + source_type=DistributedMapSourceType.INLINE, + inline_source_config=DistributedMapInlineSourceConfig( + items=serialized_items + ), + ) + + if source.inline_items is not None: + serialized_items = _build_inline_items( + source.inline_items, + source.inline_serdes or DEFAULT_JSON_SERDES, + operation_id, + durable_execution_arn, + ) + return DistributedMapSourceConfig( + source_type=DistributedMapSourceType.INLINE, + inline_source_config=DistributedMapInlineSourceConfig( + items=serialized_items + ), + max_items=source.max_items, + ) + if source.s3 is not None: + return DistributedMapSourceConfig( + source_type=DistributedMapSourceType.S3, + max_items=source.max_items, + s3_config=_build_s3_source_config(source.s3), + ) + if source.reader is not None: + reader_config = DistributedMapReaderFunctionSourceConfig( + function_name=source.reader.function_name + ) + if source.reader.initial_state is not None: + state = serialize( + serdes=source.reader.state_serdes or DEFAULT_JSON_SERDES, + value=source.reader.initial_state, + operation_id=operation_id, + durable_execution_arn=durable_execution_arn, + ) + if len(state.encode("utf-8")) > _READER_STATE_LIMIT: + msg = ( + f"reader initial_state exceeds the " + f"{_READER_STATE_LIMIT // 1024} KB limit" + ) + raise ValidationError(msg) + reader_config = DistributedMapReaderFunctionSourceConfig( + function_name=source.reader.function_name, initial_state=state + ) + return DistributedMapSourceConfig( + source_type=DistributedMapSourceType.READER_FUNCTION, + max_items=source.max_items, + reader_config=reader_config, + ) + msg = "Distributed map source has no configured items" + raise ExecutionError(msg) + + +def _build_distributed_map_options( + source: DistributedMapSource | Sequence[Any], + processor: DistributedMapProcessor, + max_concurrency: int, + config: DistributedMapConfig, + operation_id: str, + durable_execution_arn: str, +) -> DistributedMapOptions: + """Assemble the DistributedMapOptions payload from the operands and config.""" + result_collection = ( + DistributedMapResultCollectionConfig( + mode=DistributedMapResultCollectionMode.INLINE + ) + if isinstance(config, DistributedMapResultConfig) + else None + ) + return DistributedMapOptions( + max_concurrency=max_concurrency, + source=_build_source_config(source, operation_id, durable_execution_arn), + processor=_build_processor_config(processor), + destination=( + _build_destination_config(config.destination) + if config.destination is not None + else None + ), + completion_config=( + _build_completion_config(config.completion_config) + if config.completion_config is not None + else None + ), + result_collection=result_collection, + timeout_seconds=config.timeout.to_seconds() + if config.timeout is not None + else None, + ) + + +def _distributed_map_status_from_operation( + operation: Operation | None, +) -> DistributedMapStatus: + """Derive the map run status from the operation's terminal status.""" + op_status = operation.status if operation else None + resolved = ( + { + OperationStatus.SUCCEEDED: DistributedMapStatus.SUCCEEDED, + OperationStatus.FAILED: DistributedMapStatus.FAILED, + OperationStatus.STOPPED: DistributedMapStatus.STOPPED, + OperationStatus.TIMED_OUT: DistributedMapStatus.TIMED_OUT, + }.get(op_status) + if op_status is not None + else None + ) + if resolved is None: + msg = ( + "Cannot derive distributed map status from operation status " + f"{op_status.value if op_status else 'UNKNOWN'}" + ) + raise ExecutionError(msg) + return resolved + + +class DistributedMapOperationExecutor(OperationExecutor[DistributedMapSummary]): + """Executor for map run operations. + + Creates the START checkpoint if none exists, then suspends until the + backend completes the run and re-invokes the parent. On resume, a + ``DistributedMapSummary`` (or ``DistributedMapResult`` when result collection is enabled) + is built from the checkpointed ``DistributedMapDetails``. + """ + + def __init__( + self, + source: DistributedMapSource | Sequence[Any], + processor: DistributedMapProcessor, + max_concurrency: int, + state: ExecutionState, + operation_identifier: OperationIdentifier, + config: DistributedMapConfig, + ): + """Initialize the map run operation executor. + + Args: + source: The items to process (typed source or plain-list shorthand) + processor: The processor configuration + max_concurrency: Maximum concurrent processor invocations + state: The execution state + operation_identifier: The operation identifier + config: Configuration for the map run operation + """ + self.source = source + self.processor = processor + self.max_concurrency = max_concurrency + self.state = state + self.operation_identifier = operation_identifier + self.config = config + + def _resolve_summary(self, operation: Operation | None) -> DistributedMapSummary: + """Reconstruct the resolved summary/result from the terminal operation.""" + details = operation.distributed_map_details if operation else None + if details is None: + msg = "DISTRIBUTED_MAP operation succeeded but carried no DistributedMapDetails" + raise ExecutionError(msg) + status = _distributed_map_status_from_operation(operation) + completion_reason = details.completion_reason + if completion_reason is None: + msg = ( + f"DISTRIBUTED_MAP operation ended {status.value} but carried no " + f"CompletionReason" + ) + raise ExecutionError(msg) + if not isinstance(self.config, DistributedMapResultConfig): + return DistributedMapSummary( + status=status, + completion_reason=completion_reason, + success_count=details.success_count, + failure_count=details.failure_count, + unprocessed_count=details.unprocessed_count, + distributed_map_run_arn=details.distributed_map_run_arn, + completion_details=details.completion_details, + total_count=details.total_count, + ) + return DistributedMapResult( + status=status, + completion_reason=completion_reason, + success_count=details.success_count, + failure_count=details.failure_count, + unprocessed_count=details.unprocessed_count, + distributed_map_run_arn=details.distributed_map_run_arn, + completion_details=details.completion_details, + total_count=details.total_count, + all=self._deserialize_items(details, self.config.result_serdes), + ) + + def _deserialize_items( + self, details: DistributedMapDetails, result_serdes: Any + ) -> list[DistributedMapResultItem]: + """Deserialize the service result items into customer result items.""" + items: list[DistributedMapResultItem] = [] + for entry in details.results or (): + output: Any | None = None + if entry.output is not None: + output = deserialize( + serdes=result_serdes or DEFAULT_JSON_SERDES, + data=entry.output, + operation_id=self.operation_identifier.operation_id, + durable_execution_arn=self.state.durable_execution_arn, + ) + error = ( + DistributedMapItemError( + error_type=entry.error.type or "", + error_message=entry.error.message or "", + ) + if entry.error is not None + else None + ) + items.append( + DistributedMapResultItem( + item_id=entry.item_id, + status=entry.status, + output=output, + error=error, + ) + ) + return items + + def check_result_status(self) -> CheckResult[DistributedMapSummary]: + """Check operation status and create the START checkpoint if needed. + + Called twice by process() when creating synchronous checkpoints: once before + and once after, to detect if the operation completed immediately. + + Returns: + CheckResult indicating the next action to take + + Raises: + SuspendExecution: For STARTED operations waiting for completion + """ + checkpointed_result: CheckpointedResult = self.state.get_checkpoint_result( + self.operation_identifier.operation_id + ) + + # Terminal success - build the summary/result from the operation + if checkpointed_result.is_succeeded(): + operation = checkpointed_result.operation + summary = self._resolve_summary(operation) + return CheckResult.create_completed(summary) + + # Operation-level terminal failure. Every terminal state resolves with the + # summary, and throw_if_error opts into raising. + if ( + checkpointed_result.is_failed() + or checkpointed_result.is_timed_out() + or checkpointed_result.is_stopped() + ): + operation = checkpointed_result.operation + if operation is not None and operation.distributed_map_details is not None: + return CheckResult.create_completed(self._resolve_summary(operation)) + status_value = ( + checkpointed_result.status.value + if checkpointed_result.status + else "UNKNOWN" + ) + msg = ( + f"DISTRIBUTED_MAP operation ended {status_value} but carried no " + f"DistributedMapDetails" + ) + raise ExecutionError(msg) + + # Started - ready to suspend + if checkpointed_result.is_started(): + logger.debug( + "⏳ Map run %s still in progress, will suspend", + self.operation_identifier.name + or self.operation_identifier.operation_id, + ) + return CheckResult.create_is_ready_to_execute(checkpointed_result) + + # Create START checkpoint if not exists + if not checkpointed_result.is_existent(): + start_operation: OperationUpdate = ( + OperationUpdate.create_distributed_map_start( + identifier=self.operation_identifier, + distributed_map_options=_build_distributed_map_options( + source=self.source, + processor=self.processor, + max_concurrency=self.max_concurrency, + config=self.config, + operation_id=self.operation_identifier.operation_id, + durable_execution_arn=self.state.durable_execution_arn, + ), + ) + ) + # Checkpoint map run START with blocking (is_sync=True). + # Must ensure the map run is recorded before suspending execution. + self.state.create_checkpoint(operation_update=start_operation, is_sync=True) + + logger.debug( + "🚀 Map run %s started, will check for immediate completion", + self.operation_identifier.name + or self.operation_identifier.operation_id, + ) + + # Signal to process() that checkpoint was created - to recheck status + # for immediate completion before proceeding. + return CheckResult.create_started() + + # Ready to suspend (checkpoint exists but not in a terminal or started state) + return CheckResult.create_is_ready_to_execute(checkpointed_result) + + def execute( + self, _checkpointed_result: CheckpointedResult + ) -> DistributedMapSummary: + """Execute map run operation by suspending to wait for async completion. + + The map run operation doesn't execute synchronously - it suspends and + the backend runs the map run asynchronously. + + Args: + checkpointed_result: The checkpoint data (unused, but required by interface) + + Returns: + Never returns - always suspends + + Raises: + Always suspends via suspend_with_optional_resume_delay + ExecutionError: If suspend doesn't raise (should never happen) + """ + msg: str = f"Map run {self.operation_identifier.operation_id} started, suspending for completion" + suspend_with_optional_resume_delay(msg) + # This line should never be reached since suspend_with_optional_resume_delay always raises + error_msg: str = "suspend_with_optional_resume_delay should have raised an exception, but did not." + raise ExecutionError(error_msg) from None diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py index e9549d30..67d83f2f 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/plugin.py @@ -49,6 +49,7 @@ class OperationType(Enum): WAIT = "WAIT" CALLBACK = "CALLBACK" CHAINED_INVOKE = "CHAINED_INVOKE" + DISTRIBUTED_MAP = "DISTRIBUTED_MAP" def _to_invocation_status(status: ServiceInvocationStatus) -> InvocationStatus: diff --git a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py index f5ce7214..e762beb8 100644 --- a/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py +++ b/packages/aws-durable-execution-sdk-python/src/aws_durable_execution_sdk_python/state.py @@ -223,6 +223,9 @@ def create_from_operation(cls, operation: Operation) -> CheckpointedResult: result = context_details.result if context_details else None error = context_details.error if context_details else None + # OperationType.DISTRIBUTED_MAP has no operation-level result/error. + # The operation itself carries the map run's outcome. + return cls( operation=operation, status=operation.status, result=result, error=error ) diff --git a/packages/aws-durable-execution-sdk-python/tests/config_test.py b/packages/aws-durable-execution-sdk-python/tests/config_test.py index 5f230974..9b70a9b0 100644 --- a/packages/aws-durable-execution-sdk-python/tests/config_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/config_test.py @@ -22,6 +22,15 @@ WaitForConditionConfig, WaitForConditionDecision, ) +from aws_durable_execution_sdk_python.config import ( + DistributedMapCompletionConfig, + DistributedMapConfig, + S3Destination, + S3Source, + DistributedMapProcessor, + DistributedMapSource, + ReaderSourceConfig, +) def test_completion_config_defaults(): @@ -279,3 +288,165 @@ def test_completion_decision_rejects_outcome_when_not_complete(): # endregion Config validation + + +# ========================================================================== +# Distributed map config validation +# ========================================================================== + + +def test_retry_duration_out_of_range_rejected(): + with pytest.raises(ValidationError, match="between 1 minute and 6 hours"): + DistributedMapProcessor.batch("p", max_retry_duration=Duration.from_seconds(30)) + with pytest.raises(ValidationError, match="between 1 minute and 6 hours"): + DistributedMapProcessor.batch("p", max_retry_duration=Duration.from_hours(7)) + + +def test_expected_bucket_owner_must_be_12_digits(): + with pytest.raises(ValidationError, match="12-digit"): + S3Source.json_lines("s3://b/k.jsonl", expected_bucket_owner="123") + + +def test_csv_delimiter_invalid_string_rejected(): + with pytest.raises(ValidationError, match="delimiter must be one of"): + S3Source.csv("s3://b/data.csv", delimiter="BAR") + + +def test_csv_headers_duplicates_rejected(): + with pytest.raises(ValidationError, match="duplicates"): + S3Source.csv("s3://b/k.csv", headers=["a", "a"]) + + +def test_json_lines_requires_key(): + with pytest.raises(ValidationError, match="object key"): + S3Source.json_lines("s3://bucket-only") + + +def test_timeout_out_of_range_rejected(): + with pytest.raises(ValidationError, match="at most 90 days"): + DistributedMapConfig(timeout=Duration.from_days(91)) + + +def test_empty_function_name_rejected(): + with pytest.raises(ValidationError, match="non-empty"): + DistributedMapProcessor.batch("") + + +def test_valid_function_references_accepted(): + for ref in ( + "my-func", + "my-func:PROD", + "123456789012:function:my-func", + "arn:aws:lambda:us-east-1:123456789012:function:my-func", + "arn:aws:lambda:us-east-1:123456789012:function:my-func:1", + ): + # Should not raise. + DistributedMapProcessor.batch(ref) + + +def test_function_name_over_max_length_rejected(): + with pytest.raises(ValidationError, match="at most 170 characters"): + DistributedMapProcessor.batch("f" * 171) + + +def test_completion_count_and_percentage_mutually_exclusive(): + with pytest.raises(ValidationError, match="mutually exclusive"): + DistributedMapCompletionConfig( + tolerated_failure_count=1, tolerated_failure_percentage=5 + ) + + +def test_completion_sample_size_requires_percentage(): + with pytest.raises(ValidationError, match="minimum_sample_size"): + DistributedMapCompletionConfig(minimum_sample_size=10) + + +def test_completion_negative_count_rejected(): + with pytest.raises(ValidationError, match="non-negative"): + DistributedMapCompletionConfig(tolerated_failure_count=-1) + + +def test_completion_percentage_out_of_range_rejected(): + with pytest.raises(ValidationError, match="between 0 and 100"): + DistributedMapCompletionConfig(tolerated_failure_percentage=150) + + +def test_completion_sample_size_below_one_rejected(): + with pytest.raises(ValidationError, match="at least 1"): + DistributedMapCompletionConfig( + tolerated_failure_percentage=5, minimum_sample_size=0 + ) + + +def test_completion_failure_count_factory(): + assert DistributedMapCompletionConfig.failure_count(3).tolerated_failure_count == 3 + + +def test_retry_duration_below_minimum_rejected(): + with pytest.raises(ValidationError, match="1 minute and 6 hours"): + DistributedMapProcessor.batch("p", max_retry_duration=Duration.from_seconds(30)) + + +def test_max_items_below_one_rejected(): + with pytest.raises(ValidationError, match="at least 1"): + S3Source.json_lines("s3://b/k.jsonl", max_items=0) + + +def test_csv_requires_object_key(): + with pytest.raises(ValidationError, match="csv requires an S3 object key"): + S3Source.csv("s3://bucket") + + +def test_source_without_any_kind_rejected(): + with pytest.raises(ValidationError, match="exactly one of"): + DistributedMapSource() + + +def test_source_with_two_kinds_rejected(): + with pytest.raises(ValidationError, match="exactly one of"): + DistributedMapSource( + inline_items=("a",), + reader=ReaderSourceConfig(function_name="r"), + ) + + +def test_success_destination_all_false_rejected(): + with pytest.raises(ValidationError, match="success destination must include"): + S3Destination.successes( + "s3://out/ok", include_input=False, include_output=False + ) + + +def test_failure_destination_all_false_rejected(): + with pytest.raises(ValidationError, match="failure destination must include"): + S3Destination.failures("s3://out/bad", include_input=False, include_error=False) + + +def test_invalid_s3_uri_rejected(): + with pytest.raises(ValidationError, match="Invalid S3 URI"): + S3Source.json_lines("s3:///key.jsonl") + + +def test_non_s3_scheme_uri_rejected(): + with pytest.raises(ValidationError, match="must start with s3://"): + S3Source.json_lines("http://foo/bar") + + +def test_csv_empty_headers_rejected(): + with pytest.raises(ValidationError, match="must be non-empty"): + S3Source.csv("s3://b/f.csv", headers=[]) + + +def test_negative_retry_attempts_rejected(): + with pytest.raises(ValidationError, match="non-negative"): + DistributedMapProcessor.batch("p", max_retry_attempts=-5) + + +def test_batch_size_out_of_range_rejected(): + with pytest.raises(ValidationError, match="between 1 and 10000"): + DistributedMapProcessor.batch("p", batch_size=0) + + +def test_durable_execution_name_prefix_too_long_rejected(): + with pytest.raises(ValidationError, match="1 to 36 characters"): + DistributedMapProcessor.batch("p", durable_execution_name_prefix="x" * 37) diff --git a/packages/aws-durable-execution-sdk-python/tests/context_test.py b/packages/aws-durable-execution-sdk-python/tests/context_test.py index 1b75325b..6c01127f 100644 --- a/packages/aws-durable-execution-sdk-python/tests/context_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/context_test.py @@ -17,6 +17,8 @@ Duration, InvokeConfig, MapConfig, + DistributedMapConfig, + DistributedMapProcessor, ParallelBranch, ParallelConfig, StepConfig, @@ -1093,6 +1095,214 @@ def test_wait_with_time_less_than_one(mock_executor_class): # endregion wait +# region distributed_map +@pytest.mark.parametrize("max_concurrency", [0, -1, 10001, 50000]) +def test_map_run_rejects_out_of_range_max_concurrency(max_concurrency: int): + """Test distributed_map rejects a max_concurrency outside the service range.""" + context = create_test_context() + + with pytest.raises( + ValidationError, match="max_concurrency must be between 1 and 10000" + ): + context.distributed_map( + ["a"], + DistributedMapProcessor.batch("test_processor"), + max_concurrency=max_concurrency, + ) + + +@pytest.mark.parametrize( + "source", + ["s3://bucket", b"bytes", bytearray(b"ba"), {"a": 1}, {"a", "b"}], +) +def test_distributed_map_rejects_non_list_source(source): + """Test distributed_map rejects sources that are not a DistributedMapSource or list/tuple.""" + context = create_test_context() + + with pytest.raises(ValidationError, match="list/tuple"): + context.distributed_map( + source, + DistributedMapProcessor.batch("test_processor"), + max_concurrency=1, + ) + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_basic(mock_executor_class): + """distributed_map builds a DISTRIBUTED_MAP executor and returns its process() result.""" + mock_executor = MagicMock() + mock_executor.process.return_value = "map_summary" + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + + context = create_test_context(state=mock_state) + expected_operation_id = next(operation_id_sequence()) + + result = context.distributed_map(["a", "b"], processor, max_concurrency=4) + + assert result == "map_summary" + mock_executor_class.assert_called_once_with( + state=mock_state, + operation_identifier=OperationIdentifier( + expected_operation_id, OperationSubType.DISTRIBUTED_MAP, None, None + ), + source=["a", "b"], + processor=processor, + max_concurrency=4, + config=ANY, + ) + mock_executor.process.assert_called_once() + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_with_name_and_config(mock_executor_class): + """distributed_map forwards name into the identifier and passes the given config through.""" + mock_executor = MagicMock() + mock_executor.process.return_value = "configured_summary" + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + config = DistributedMapConfig() + + context = create_test_context(state=mock_state) + [context._create_step_id() for _ in range(5)] # Set counter to 5 # noqa: SLF001 + + result = context.distributed_map( + ["a"], processor, max_concurrency=2, name="named_map", config=config + ) + + seq = operation_id_sequence() + [next(seq) for _ in range(5)] + expected_id = next(seq) + + assert result == "configured_summary" + mock_executor_class.assert_called_once_with( + state=mock_state, + operation_identifier=OperationIdentifier( + expected_id, OperationSubType.DISTRIBUTED_MAP, None, "named_map" + ), + source=["a"], + processor=processor, + max_concurrency=2, + config=config, + ) + mock_executor.process.assert_called_once() + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_with_parent_id(mock_executor_class): + """distributed_map propagates the parent_id into the operation identifier.""" + mock_executor = MagicMock() + mock_executor.process.return_value = "parent_summary" + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + + context = create_test_context(state=mock_state, parent_id="parent123") + [context._create_step_id() for _ in range(2)] # Set counter to 2 # noqa: SLF001 + + context.distributed_map(["a"], processor, max_concurrency=1) + + seq = operation_id_sequence("parent123") + [next(seq) for _ in range(2)] + expected_id = next(seq) + + mock_executor_class.assert_called_once_with( + state=mock_state, + operation_identifier=OperationIdentifier( + expected_id, OperationSubType.DISTRIBUTED_MAP, "parent123", None + ), + source=["a"], + processor=processor, + max_concurrency=1, + config=ANY, + ) + mock_executor.process.assert_called_once() + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_increments_counter(mock_executor_class): + """distributed_map increments the step counter once per call.""" + mock_executor = MagicMock() + mock_executor.process.return_value = "result" + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + + context = create_test_context(state=mock_state) + [context._create_step_id() for _ in range(10)] # Set counter to 10 # noqa: SLF001 + + context.distributed_map(["a"], processor, max_concurrency=1) + context.distributed_map(["b"], processor, max_concurrency=1) + + seq = operation_id_sequence() + [next(seq) for _ in range(10)] + expected_id1 = next(seq) + expected_id2 = next(seq) + + assert context._step_counter.get_current() == 12 # noqa: SLF001 + assert mock_executor_class.call_args_list[0][1][ + "operation_identifier" + ] == OperationIdentifier(expected_id1, OperationSubType.DISTRIBUTED_MAP, None, None) + assert mock_executor_class.call_args_list[1][1][ + "operation_identifier" + ] == OperationIdentifier(expected_id2, OperationSubType.DISTRIBUTED_MAP, None, None) + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_defaults_config_when_none(mock_executor_class): + """distributed_map builds a default DistributedMapConfig when none is given.""" + mock_executor = MagicMock() + mock_executor.process.return_value = "summary" + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + + context = create_test_context(state=mock_state) + context.distributed_map(["a"], processor, max_concurrency=1) + + passed_config = mock_executor_class.call_args[1]["config"] + assert isinstance(passed_config, DistributedMapConfig) + + +@patch("aws_durable_execution_sdk_python.context.DistributedMapOperationExecutor") +def test_distributed_map_returns_process_result(mock_executor_class): + """distributed_map returns whatever executor.process() returns (summary or result).""" + mock_executor = MagicMock() + sentinel = object() + mock_executor.process.return_value = sentinel + mock_executor_class.return_value = mock_executor + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = ( + "arn:aws:durable:us-east-1:123456789012:execution/test" + ) + processor = DistributedMapProcessor.batch("test_processor") + + context = create_test_context(state=mock_state) + result = context.distributed_map(["a"], processor, max_concurrency=1) + + assert result is sentinel + + +# endregion distributed_map + + # region run_in_child_context @patch("aws_durable_execution_sdk_python.context.child_handler") def test_run_in_child_context_basic(mock_handler): diff --git a/packages/aws-durable-execution-sdk-python/tests/dmap/handlers_test.py b/packages/aws-durable-execution-sdk-python/tests/dmap/handlers_test.py new file mode 100644 index 00000000..1f467594 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/dmap/handlers_test.py @@ -0,0 +1,253 @@ +"""Unit tests for the non-durable distributed map authoring wrappers.""" + +from __future__ import annotations + +import threading +import time +from unittest.mock import MagicMock + +import pytest + +from aws_durable_execution_sdk_python.config import ( + DistributedMapConfig, + DistributedMapProcessor, + DistributedMapStatus, +) +from aws_durable_execution_sdk_python.dmap.handlers import ( + ReaderPage, + distributed_map_batch_handler, + distributed_map_item_handler, + durable_distributed_map_item_handler, + distributed_map_reader, +) +from aws_durable_execution_sdk_python.exceptions import ExecutionError, ValidationError +from aws_durable_execution_sdk_python.lambda_service import ( + DistributedMapCompletionReason, + DistributedMapDetails, + Operation, + OperationStatus, + OperationType, +) +from aws_durable_execution_sdk_python.operation.dmap import ( + DistributedMapOperationExecutor, + _distributed_map_status_from_operation, +) + + +def _terminal_operation( + status: OperationStatus, details: DistributedMapDetails +) -> Operation: + return Operation( + operation_id="dmap-id", + operation_type=OperationType.DISTRIBUTED_MAP, + status=status, + distributed_map_details=details, + ) + + +def test_status_derived_from_operation_when_details_omit_status(): + """Status comes from the operation, not from details.""" + details = DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=3, + total_count=3, + ) + op = _terminal_operation(OperationStatus.SUCCEEDED, details) + assert _distributed_map_status_from_operation(op) is DistributedMapStatus.SUCCEEDED + + +def test_status_from_operation_rejects_non_terminal(): + """A non-terminal operation status cannot map to a DistributedMapStatus.""" + op = _terminal_operation( + OperationStatus.STARTED, + DistributedMapDetails(success_count=0), + ) + with pytest.raises(ExecutionError, match="Cannot derive distributed map status"): + _distributed_map_status_from_operation(op) + + +def test_resolved_summary_uses_operation_status_not_details(): + """_resolve_summary carries the derived status when details omit it.""" + details = DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=2, + failure_count=0, + unprocessed_count=0, + total_count=2, + ) + executor = DistributedMapOperationExecutor( + source=["a", "b"], + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + state=MagicMock(), + operation_identifier=MagicMock(), + config=DistributedMapConfig(), + ) + summary = executor._resolve_summary( + _terminal_operation(OperationStatus.SUCCEEDED, details) + ) + assert summary.status is DistributedMapStatus.SUCCEEDED + assert summary.completion_reason is DistributedMapCompletionReason.ALL_COMPLETED + assert summary.success_count == 2 + + +def test_item_handler_invalid_report_rejected(): + with pytest.raises(ValidationError, match="report must be"): + distributed_map_item_handler(lambda x: x, report="bogus") + + +def test_durable_item_handler_invalid_report_rejected(): + with pytest.raises(ValidationError, match="report must be"): + durable_distributed_map_item_handler(lambda _ctx, item: item, report="bogus") + + +def test_item_handler_rejects_concurrency_below_one(): + with pytest.raises(ValidationError, match="concurrency must be at least 1"): + distributed_map_item_handler(lambda x: x, concurrency=0) + with pytest.raises(ValidationError, match="concurrency must be at least 1"): + distributed_map_item_handler(lambda x: x, concurrency=-1) + + +def test_item_handler_runs_items_one_at_a_time_by_default(): + lock = threading.Lock() + in_flight = 0 + peak = 0 + + def process(x): + nonlocal in_flight, peak + with lock: + in_flight += 1 + peak = max(peak, in_flight) + # Hold the worker so that a wider pool would overlap this item with the + # next one, which is what makes the peak below meaningful. + time.sleep(0.02) + with lock: + in_flight -= 1 + return x + + handler = distributed_map_item_handler(process) + handler({"records": [{"itemId": str(i), "body": "1"} for i in range(4)]}) + assert peak == 1 + + +def test_item_handler_honors_requested_concurrency(): + # Each item blocks until all four have arrived, so the batch only finishes + # if the pool really runs four items at once. At a smaller pool size the + # first item waits out the timeout and every item fails. + barrier = threading.Barrier(4, timeout=5) + + def process(x): + barrier.wait() + return x + + handler = distributed_map_item_handler(process, concurrency=4) + resp = handler({"records": [{"itemId": str(i), "body": str(i)} for i in range(4)]}) + assert resp["batchItemFailures"] == [] + # The four items are released together, so completion order is arbitrary + # while the reported results stay in item order. + assert resp["batchItemResults"] == [ + {"itemIdentifier": "0", "output": "0"}, + {"itemIdentifier": "1", "output": "1"}, + {"itemIdentifier": "2", "output": "2"}, + {"itemIdentifier": "3", "output": "3"}, + ] + + +def test_item_handler_reports_results_in_order(): + handler = distributed_map_item_handler(lambda x: x * 2) + resp = handler( + {"records": [{"itemId": "0", "body": "2"}, {"itemId": "1", "body": "3"}]} + ) + assert resp["batchItemResults"] == [ + {"itemIdentifier": "0", "output": "4"}, + {"itemIdentifier": "1", "output": "6"}, + ] + assert resp["batchItemFailures"] == [] + + +def test_item_handler_captures_failures(): + def process(x): + if x == "bad": + msg = "boom" + raise ValueError(msg) + return x + + handler = distributed_map_item_handler(process) + resp = handler( + {"records": [{"itemId": "0", "body": '"ok"'}, {"itemId": "1", "body": '"bad"'}]} + ) + assert resp["batchItemResults"] == [{"itemIdentifier": "0", "output": '"ok"'}] + assert resp["batchItemFailures"] == [ + { + "itemIdentifier": "1", + "error": {"errorType": "ValueError", "errorMessage": "boom"}, + } + ] + + +def test_item_handler_failures_form_reports_only_failures(): + handler = distributed_map_item_handler(lambda x: x, report="failures") + resp = handler({"records": [{"itemId": "0", "body": "1"}]}) + assert resp == {"batchItemFailures": []} + + +def test_batch_handler_success_and_propagates_error(): + seen: list = [] + handler = distributed_map_batch_handler(lambda items: seen.extend(items) or "done") + assert ( + handler( + {"records": [{"itemId": "0", "body": "1"}, {"itemId": "1", "body": "2"}]} + ) + == "done" + ) + assert seen == [1, 2] + + def boom(_items): + msg = "batch failed" + raise RuntimeError(msg) + + failing = distributed_map_batch_handler(boom) + with pytest.raises(RuntimeError, match="batch failed"): + failing({"records": [{"itemId": "0", "body": "1"}]}) + + +def test_reader_returns_items_and_next_state_then_exhausts(): + def read(state): + if state is None: + return ReaderPage(items=[1, 2], next_state={"page": 1}) + return ReaderPage(items=[3]) + + handler = distributed_map_reader(read) + first = handler({"state": None, "maxItems": 10}) + assert first["items"] == [1, 2] + assert first["nextState"] == '{"page": 1}' + + second = handler({"state": '{"page": 1}', "maxItems": 10}) + assert second["items"] == [3] + assert "nextState" not in second + + +def test_reader_rejects_page_over_max_items(): + handler = distributed_map_reader(lambda _s: ReaderPage(items=[1, 2, 3])) + with pytest.raises(ValidationError, match="exceeding maxItems"): + handler({"state": None, "maxItems": 2}) + + +def test_reader_rejects_oversized_next_state(): + handler = distributed_map_reader( + lambda _s: ReaderPage(items=[1], next_state="x" * 40_000) + ) + with pytest.raises(ValidationError, match="32 KB limit"): + handler({"state": None, "maxItems": 10}) + + +def test_item_handler_rejects_non_processor_envelope(): + handler = distributed_map_item_handler(lambda x: x) + with pytest.raises(ValidationError, match="processor envelope"): + handler({}) + + +def test_reader_rejects_non_reader_envelope(): + handler = distributed_map_reader(lambda _state: ReaderPage(items=[])) + with pytest.raises(ValidationError, match="reader envelope"): + handler({"state": None}) diff --git a/packages/aws-durable-execution-sdk-python/tests/dmap/models_test.py b/packages/aws-durable-execution-sdk-python/tests/dmap/models_test.py new file mode 100644 index 00000000..9048ae46 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/dmap/models_test.py @@ -0,0 +1,173 @@ +"""Unit tests for the distributed map result types.""" + +from __future__ import annotations + +import pytest + +from aws_durable_execution_sdk_python.config import ( + DistributedMapCompletionReason, + DistributedMapItemStatus, + DistributedMapStatus, +) +from aws_durable_execution_sdk_python.dmap.models import ( + DistributedMapItemError, + DistributedMapResult, + DistributedMapResultItem, + DistributedMapSummary, +) +from aws_durable_execution_sdk_python.exceptions import DistributedMapError + + +def _result(status, completion_reason, failure_count, items): + return DistributedMapResult( + status=status, + completion_reason=completion_reason, + success_count=len(items) - failure_count, + failure_count=failure_count, + unprocessed_count=0, + all=items, + ) + + +def test_summary_throw_if_error(): + ok = DistributedMapSummary( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=2, + failure_count=0, + unprocessed_count=0, + ) + ok.throw_if_error() # no raise + + failed = DistributedMapSummary( + status=DistributedMapStatus.FAILED, + completion_reason=DistributedMapCompletionReason.FAILURE_TOLERANCE_EXCEEDED, + success_count=0, + failure_count=1, + unprocessed_count=0, + ) + with pytest.raises(DistributedMapError): + failed.throw_if_error() + + succeeded_with_failures = DistributedMapSummary( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=1, + unprocessed_count=0, + ) + with pytest.raises(DistributedMapError): + succeeded_with_failures.throw_if_error() + + +def test_map_run_result_succeeded_failed_filters(): + items = [ + DistributedMapResultItem( + item_id="0", status=DistributedMapItemStatus.SUCCEEDED, output=1 + ), + DistributedMapResultItem( + item_id="1", + status=DistributedMapItemStatus.FAILED, + error=DistributedMapItemError(error_type="E", error_message="boom"), + ), + ] + result = DistributedMapResult( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=1, + unprocessed_count=0, + all=items, + ) + assert [i.item_id for i in result.succeeded()] == ["0"] + assert [i.item_id for i in result.failed()] == ["1"] + assert result.get_results() == [1] + assert result.get_errors()[0].error_message == "boom" + + +def test_summary_has_failure_false_without_failures(): + summary = DistributedMapSummary( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=0, + failure_count=0, + unprocessed_count=0, + ) + assert summary.has_failure is False + + +def test_result_throw_if_error_raises_first_item_error(): + result = _result( + DistributedMapStatus.SUCCEEDED, + DistributedMapCompletionReason.ALL_COMPLETED, + 1, + [ + DistributedMapResultItem( + item_id="0", status=DistributedMapItemStatus.SUCCEEDED, output=1 + ), + DistributedMapResultItem( + item_id="1", + status=DistributedMapItemStatus.FAILED, + error=DistributedMapItemError(error_type="E", error_message="boom"), + ), + ], + ) + with pytest.raises(DistributedMapError, match="E: boom"): + result.throw_if_error() + + +def test_result_throw_if_error_names_item_without_detail(): + result = _result( + DistributedMapStatus.SUCCEEDED, + DistributedMapCompletionReason.ALL_COMPLETED, + 1, + [ + DistributedMapResultItem( + item_id="7", status=DistributedMapItemStatus.FAILED, error=None + ) + ], + ) + with pytest.raises(DistributedMapError, match="item 7 failed"): + result.throw_if_error() + + +def test_result_throw_if_error_falls_back_to_summary_when_no_items(): + result = DistributedMapResult( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=0, + failure_count=2, + unprocessed_count=0, + all=[], + ) + with pytest.raises(DistributedMapError, match="2 item"): + result.throw_if_error() + + +def test_result_throw_if_error_run_level_failure(): + result = DistributedMapResult( + status=DistributedMapStatus.FAILED, + completion_reason=DistributedMapCompletionReason.FAILURE_TOLERANCE_EXCEEDED, + success_count=0, + failure_count=1, + unprocessed_count=0, + all=[], + ) + with pytest.raises(DistributedMapError, match="Map run ended FAILED"): + result.throw_if_error() + + +def test_result_throw_if_error_clean_success_does_not_raise(): + result = DistributedMapResult( + status=DistributedMapStatus.SUCCEEDED, + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=0, + unprocessed_count=0, + all=[ + DistributedMapResultItem( + item_id="0", status=DistributedMapItemStatus.SUCCEEDED, output=1 + ) + ], + ) + result.throw_if_error() diff --git a/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_helpers_int_test.py b/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_helpers_int_test.py new file mode 100644 index 00000000..f6336d61 --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_helpers_int_test.py @@ -0,0 +1,146 @@ +"""Integration tests for the durable distributed map authoring wrappers. + +The durable variants are ``@durable_execution`` handlers; they run through the +full invocation harness with a mocked checkpoint backend. +""" + +from __future__ import annotations + +import json +from unittest.mock import Mock, patch + +from aws_durable_execution_sdk_python.dmap.handlers import ( + durable_distributed_map_batch_handler, + durable_distributed_map_item_handler, +) +from aws_durable_execution_sdk_python.execution import ( + InvocationStatus, + durable_execution, # noqa: F401 (ensures decorator import path is valid) +) +from aws_durable_execution_sdk_python.lambda_service import ( + CheckpointOutput, + CheckpointUpdatedExecutionState, + Operation, + OperationStatus, + OperationType, +) + +_ARN = "arn:aws:lambda:us-east-1:123456789012:function:proc:1/durable-execution/execution-1" + + +def _lambda_context(): + ctx = Mock() + ctx.aws_request_id = "test-request-id" + ctx.client_context = None + ctx.identity = None + ctx._epoch_deadline_time_in_ms = 0 # noqa: SLF001 + ctx.invoked_function_arn = "test-arn" + ctx.tenant_id = None + return ctx + + +def _event(records: list[dict]): + return { + "DurableExecutionArn": _ARN, + "CheckpointToken": "test-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "execution-1", + "Type": "EXECUTION", + "Status": "STARTED", + "ExecutionDetails": { + "InputPayload": json.dumps({"records": records}) + }, + } + ], + "NextMarker": "", + }, + "LocalRunner": True, + } + + +def _run(handler, records: list[dict]): + operations = [ + Operation( + operation_id="execution-1", + operation_type=OperationType.EXECUTION, + status=OperationStatus.STARTED, + ) + ] + + def mock_checkpoint( + durable_execution_arn, checkpoint_token, updates, client_token="token" + ): # noqa: S107 + for update in updates: + operations.append( + Operation( + operation_id=update.operation_id, + operation_type=update.operation_type, + status=OperationStatus.STARTED, + parent_id=update.parent_id, + ) + ) + return CheckpointOutput( + checkpoint_token="new_token", # noqa: S106 + new_execution_state=CheckpointUpdatedExecutionState( + operations=operations.copy() + ), + ) + + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient" + ) as mock_client_class: + mock_client = Mock() + mock_client_class.initialize_client.return_value = mock_client + mock_client.checkpoint = mock_checkpoint + return handler(_event(records), _lambda_context()) + + +def test_durable_item_handler_reports_results(): + handler = durable_distributed_map_item_handler(lambda _ctx, item: item * 2) + result = _run(handler, [{"itemId": "0", "body": "2"}, {"itemId": "1", "body": "3"}]) + assert result["Status"] == InvocationStatus.SUCCEEDED.value + data = json.loads(result["Result"]) + assert data["batchItemResults"] == [ + {"itemIdentifier": "0", "output": "4"}, + {"itemIdentifier": "1", "output": "6"}, + ] + assert data["batchItemFailures"] == [] + + +def test_durable_item_handler_captures_failure(): + def process(_ctx, item): + if item == "bad": + msg = "boom" + raise ValueError(msg) + return item + + handler = durable_distributed_map_item_handler(process) + result = _run( + handler, [{"itemId": "0", "body": '"ok"'}, {"itemId": "1", "body": '"bad"'}] + ) + assert result["Status"] == InvocationStatus.SUCCEEDED.value + data = json.loads(result["Result"]) + assert data["batchItemResults"] == [{"itemIdentifier": "0", "output": '"ok"'}] + assert data["batchItemFailures"][0]["itemIdentifier"] == "1" + assert data["batchItemFailures"][0]["error"]["errorType"] == "ValueError" + + +def test_durable_batch_handler_returns_value(): + handler = durable_distributed_map_batch_handler( + lambda _ctx, items: {"count": len(items)} + ) + result = _run(handler, [{"itemId": "0", "body": "1"}, {"itemId": "1", "body": "2"}]) + assert result["Status"] == InvocationStatus.SUCCEEDED.value + assert json.loads(result["Result"]) == {"count": 2} + + +def test_durable_item_handler_failures_form(): + handler = durable_distributed_map_item_handler( + lambda _ctx, item: item, report="failures" + ) + result = _run(handler, [{"itemId": "0", "body": "1"}]) + assert result["Status"] == InvocationStatus.SUCCEEDED.value + data = json.loads(result["Result"]) + assert data == {"batchItemFailures": []} diff --git a/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_int_test.py b/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_int_test.py new file mode 100644 index 00000000..b07917bb --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/e2e/dmap_int_test.py @@ -0,0 +1,246 @@ +"""Integration tests for ctx.distributed_map through a full durable_execution invocation. + +Drives the real suspend/resume flow: the first invocation starts the DISTRIBUTED_MAP +operation and suspends (PENDING); a replay invocation with the operation +completed (carrying DistributedMapDetails) resolves to the summary/result. The backend +is mocked; no emulator or real service is involved. +""" + +from __future__ import annotations + +import json +from typing import Any +from unittest.mock import Mock, patch + +from aws_durable_execution_sdk_python.dmap.models import DistributedMapResult +from aws_durable_execution_sdk_python.config import ( + DistributedMapResultConfig, + DistributedMapProcessor, +) +from aws_durable_execution_sdk_python.context import DurableContext +from aws_durable_execution_sdk_python.execution import ( + InvocationStatus, + durable_execution, +) +from aws_durable_execution_sdk_python.lambda_service import ( + CheckpointOutput, + CheckpointUpdatedExecutionState, + Operation, + OperationStatus, + OperationType, +) +from tests.test_helpers import operation_id_sequence + + +def _lambda_context(): + ctx = Mock() + ctx.aws_request_id = "test-request-id" + ctx.client_context = None + ctx.identity = None + ctx._epoch_deadline_time_in_ms = 0 # noqa: SLF001 + ctx.invoked_function_arn = "test-arn" + ctx.tenant_id = None + return ctx + + +_ARN = ( + "arn:aws:lambda:us-east-1:123456789012:function:test-func:1" + "/durable-execution/exec-001/inv-001" +) + + +def _initial_event(): + return { + "DurableExecutionArn": _ARN, + "CheckpointToken": "test-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "execution-1", + "Type": "EXECUTION", + "Status": "STARTED", + "ExecutionDetails": {"InputPayload": "{}"}, + } + ], + "NextMarker": "", + }, + "LocalRunner": True, + } + + +def _replay_event(distributed_map_details: dict): + distributed_map_id = next(operation_id_sequence()) + return distributed_map_id, { + "DurableExecutionArn": _ARN, + "CheckpointToken": "test-token", + "InitialExecutionState": { + "Operations": [ + { + "Id": "execution-1", + "Type": "EXECUTION", + "Status": "STARTED", + "ExecutionDetails": {"InputPayload": "{}"}, + }, + { + "Id": distributed_map_id, + "Type": "DISTRIBUTED_MAP", + "SubType": "DistributedMap", + "Status": "SUCCEEDED", + "DistributedMapDetails": distributed_map_details, + }, + ], + "NextMarker": "", + }, + "LocalRunner": True, + } + + +def _tracking_checkpoint(): + """Checkpoint mock that records created operations as STARTED.""" + calls: list = [] + operations = [ + Operation( + operation_id="execution-1", + operation_type=OperationType.EXECUTION, + status=OperationStatus.STARTED, + ) + ] + + def mock_checkpoint( + durable_execution_arn, checkpoint_token, updates, client_token="token" + ): # noqa: S107 + calls.append(updates) + for update in updates: + operations.append( + Operation( + operation_id=update.operation_id, + operation_type=update.operation_type, + status=OperationStatus.STARTED, + parent_id=update.parent_id, + ) + ) + return CheckpointOutput( + checkpoint_token="new_token", # noqa: S106 + new_execution_state=CheckpointUpdatedExecutionState( + operations=operations.copy() + ), + ) + + return calls, mock_checkpoint + + +def _run(handler, event): + with patch( + "aws_durable_execution_sdk_python.execution.LambdaClient" + ) as mock_client_class: + mock_client = Mock() + mock_client_class.initialize_client.return_value = mock_client + _calls, mock_checkpoint = _tracking_checkpoint() + mock_client.checkpoint = mock_checkpoint + return handler(event, _lambda_context()) + + +def test_map_run_suspends_then_resumes_with_summary(): + @durable_execution + def handler(event, context: DurableContext) -> dict[str, Any]: + summary = context.distributed_map( + ["a", "b"], + DistributedMapProcessor.batch("proc"), + max_concurrency=2, + ) + return { + "status": summary.status.value, + "success": summary.success_count, + "failure": summary.failure_count, + "distributed_map_run_arn": summary.distributed_map_run_arn, + } + + # First invocation suspends. + first = _run(handler, _initial_event()) + assert first["Status"] == InvocationStatus.PENDING.value + + # Replay with the run completed. + _map_run_id, replay_event = _replay_event( + { + "Status": "SUCCEEDED", + "CompletionReason": "ALL_COMPLETED", + "SuccessCount": 2, + "FailureCount": 0, + "UnprocessedCount": 0, + "TotalCount": 2, + "DistributedMapRunArn": "arn:aws:lambda:us-east-1:123456789012:function:fn:$LATEST/durable-execution/exec1/invoke1/distributed-map-run/abc", + } + ) + replay = _run(handler, replay_event) + assert replay["Status"] == InvocationStatus.SUCCEEDED.value + data = json.loads(replay["Result"]) + assert data == { + "status": "SUCCEEDED", + "success": 2, + "failure": 0, + "distributed_map_run_arn": "arn:aws:lambda:us-east-1:123456789012:function:fn:$LATEST/durable-execution/exec1/invoke1/distributed-map-run/abc", + } + + +def test_map_run_result_config_returns_items(): + @durable_execution + def handler(event, context: DurableContext) -> dict[str, Any]: + result = context.distributed_map( + ["a", "b"], + DistributedMapProcessor.item_results("proc"), + max_concurrency=2, + config=DistributedMapResultConfig(), + ) + assert isinstance(result, DistributedMapResult) + return { + "results": result.get_results(), + "errors": [e.error_message for e in result.get_errors()], + } + + _map_run_id, replay_event = _replay_event( + { + "Status": "SUCCEEDED", + "CompletionReason": "ALL_COMPLETED", + "SuccessCount": 1, + "FailureCount": 1, + "UnprocessedCount": 0, + "TotalCount": 2, + "Results": [ + {"ItemId": "0", "Status": "SUCCEEDED", "Output": "5"}, + { + "ItemId": "1", + "Status": "FAILED", + "Error": {"ErrorType": "E", "ErrorMessage": "boom"}, + }, + ], + } + ) + replay = _run(handler, replay_event) + assert replay["Status"] == InvocationStatus.SUCCEEDED.value + data = json.loads(replay["Result"]) + assert data == {"results": [5], "errors": ["boom"]} + + +def test_map_run_throw_if_error_fails_execution(): + @durable_execution + def handler(event, context: DurableContext) -> dict[str, Any]: + summary = context.distributed_map( + ["a"], + DistributedMapProcessor.batch("proc"), + max_concurrency=1, + ) + summary.throw_if_error() + return {"status": summary.status.value} + + _map_run_id, replay_event = _replay_event( + { + "Status": "FAILED", + "CompletionReason": "FAILURE_TOLERANCE_EXCEEDED", + "SuccessCount": 0, + "FailureCount": 1, + "UnprocessedCount": 0, + "TotalCount": 1, + } + ) + replay = _run(handler, replay_event) + assert replay["Status"] == InvocationStatus.FAILED.value diff --git a/packages/aws-durable-execution-sdk-python/tests/exceptions_test.py b/packages/aws-durable-execution-sdk-python/tests/exceptions_test.py index 540f09fa..9f1a906e 100644 --- a/packages/aws-durable-execution-sdk-python/tests/exceptions_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/exceptions_test.py @@ -16,6 +16,7 @@ ChildContextError, CheckpointError, CheckpointErrorCategory, + DistributedMapError, DurableApiErrorCategory, DurableExecutionsError, DurableOperationError, @@ -257,6 +258,7 @@ def test_durable_operation_error_with_none_message(): [ StepError, InvokeError, + DistributedMapError, ChildContextError, WaitForConditionError, CallbackError, diff --git a/packages/aws-durable-execution-sdk-python/tests/lambda_service_test.py b/packages/aws-durable-execution-sdk-python/tests/lambda_service_test.py index 516930e2..aa5a10e0 100644 --- a/packages/aws-durable-execution-sdk-python/tests/lambda_service_test.py +++ b/packages/aws-durable-execution-sdk-python/tests/lambda_service_test.py @@ -42,6 +42,10 @@ WaitOptions, _is_in_var_dir, ) +from aws_durable_execution_sdk_python.lambda_service import ( + DistributedMapDetails, + DistributedMapResultItem, +) # ============================================================================= @@ -1271,6 +1275,41 @@ def test_operation_from_dict_no_options(): assert operation.operation_id == "test-id" +def test_operation_from_dict_in_flight_distributed_map_details_parsed_without_status(): + """In-flight DistributedMapDetails (no Status/CompletionReason) parses without raising, leaving both None.""" + data = { + "Id": "dmap-id", + "Type": "DISTRIBUTED_MAP", + "Status": "STARTED", + "DistributedMapDetails": { + "TotalCount": 3, + "SuccessCount": 1, + }, + } + operation = Operation.from_dict(data) + assert operation.distributed_map_details is not None + assert operation.distributed_map_details.completion_reason is None + assert operation.distributed_map_details.total_count == 3 + assert operation.distributed_map_details.success_count == 1 + + +def test_operation_from_dict_terminal_distributed_map_details_parsed(): + """Terminal DistributedMapDetails parses into distributed_map_details.""" + data = { + "Id": "dmap-id", + "Type": "DISTRIBUTED_MAP", + "Status": "SUCCEEDED", + "DistributedMapDetails": { + "CompletionReason": "ALL_COMPLETED", + "TotalCount": 3, + "SuccessCount": 3, + }, + } + operation = Operation.from_dict(data) + assert operation.distributed_map_details is not None + assert operation.distributed_map_details.completion_reason.value == "ALL_COMPLETED" + + def test_operation_from_dict_individual_options(): """Test Operation.from_dict with each option type individually.""" # Test with just ContextOptions @@ -3029,3 +3068,60 @@ def test_operation_to_dict_omits_absent_context_error_and_replay_children(): context = op.to_dict()["ContextDetails"] assert context == {"Result": '"hello"'} + + +# ========================================================================== +# Distributed map wire round-trips +# ========================================================================== + + +def test_map_run_details_from_dict_parses_results(): + data = { + "Status": "SUCCEEDED", + "CompletionReason": "ALL_COMPLETED", + "SuccessCount": 1, + "FailureCount": 1, + "UnprocessedCount": 0, + "TotalCount": 2, + "DistributedMapRunArn": "arn:aws:lambda:us-east-1:123456789012:map-run:x", + "Results": [ + {"ItemId": "0", "Status": "SUCCEEDED", "Output": 5}, + { + "ItemId": "1", + "Status": "FAILED", + "Error": {"ErrorType": "E", "ErrorMessage": "boom"}, + }, + ], + } + details = DistributedMapDetails.from_dict(data) + assert details.total_count == 2 + assert details.results is not None + assert details.results[0].item_id == "0" + assert details.results[1].error is not None + assert details.results[1].error.type == "E" + + +def test_details_missing_completion_reason_parses_with_none(): + # In-flight runs omit CompletionReason too; parsing must succeed with None. + details = DistributedMapDetails.from_dict({"SuccessCount": 2}) + assert details.completion_reason is None + assert details.success_count == 2 + + +def test_result_item_wire_round_trip(): + item = DistributedMapResultItem.from_dict( + {"ItemId": "0", "Status": "SUCCEEDED", "Output": {"x": 1}} + ) + assert item.output == {"x": 1} + assert item.to_dict() == {"ItemId": "0", "Status": "SUCCEEDED", "Output": {"x": 1}} + + failed = DistributedMapResultItem.from_dict( + { + "ItemId": "1", + "Status": "FAILED", + "Error": {"ErrorType": "E", "ErrorMessage": "boom"}, + } + ) + dumped = failed.to_dict() + assert dumped["Error"]["ErrorType"] == "E" + assert "Output" not in dumped diff --git a/packages/aws-durable-execution-sdk-python/tests/operation/dmap_test.py b/packages/aws-durable-execution-sdk-python/tests/operation/dmap_test.py new file mode 100644 index 00000000..305a507b --- /dev/null +++ b/packages/aws-durable-execution-sdk-python/tests/operation/dmap_test.py @@ -0,0 +1,1388 @@ +"""Unit tests for map run handler.""" + +from __future__ import annotations + +import json +from unittest.mock import Mock, patch + +import pytest + +from aws_durable_execution_sdk_python.dmap.models import ( + DistributedMapResult, + DistributedMapSummary, +) +from aws_durable_execution_sdk_python.config import ( + Duration, + DistributedMapCompletionReason, + DistributedMapItemStatus, + DistributedMapStatus, + DistributedMapCompletionConfig, + DistributedMapConfig, + DistributedMapResultConfig, + DistributedMapCsvDelimiter, + DistributedMapDestinationConfig, + InlineSource, + ReaderSource, + S3Destination, + S3Source, + DistributedMapProcessor, +) +from aws_durable_execution_sdk_python.exceptions import ( + ExecutionError, + DistributedMapError, + SuspendExecution, + ValidationError, +) +from aws_durable_execution_sdk_python.identifier import OperationIdentifier +from aws_durable_execution_sdk_python.lambda_service import ( + ErrorObject, + DistributedMapDetails, + DistributedMapFunctionResponseType, + DistributedMapOptions, + DistributedMapResultCollectionMode, + DistributedMapResultItem as DistributedMapResultItemApi, + DistributedMapSourceType, + Operation, + OperationAction, + OperationStatus, + OperationSubType, + OperationType, +) +from aws_durable_execution_sdk_python.operation.dmap import ( + DistributedMapOperationExecutor, +) +from aws_durable_execution_sdk_python.state import CheckpointedResult, ExecutionState + + +# Test helper - wraps DistributedMapOperationExecutor with a simple handler signature. +def distributed_map_handler( + source, processor, max_concurrency, state, operation_identifier, config=None +): + """Test helper that wraps DistributedMapOperationExecutor and runs it. + + ``processor`` may be a function-name string (wrapped as a batch-outcome + processor) or an already-built DistributedMapProcessor. + """ + if not config: + config = DistributedMapConfig() + if isinstance(processor, str): + processor = DistributedMapProcessor.batch(processor) + executor = DistributedMapOperationExecutor( + source=source, + processor=processor, + max_concurrency=max_concurrency, + state=state, + operation_identifier=operation_identifier, + config=config, + ) + return executor.process() + + +def _identifier( + operation_id: str, name: str | None = "test_map_run" +) -> OperationIdentifier: + return OperationIdentifier( + operation_id, OperationSubType.DISTRIBUTED_MAP, None, name + ) + + +def test_map_run_handler_already_succeeded(): + """Test distributed_map_handler returns a summary when the operation already succeeded.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + operation = Operation( + operation_id="mr1", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=5, + failure_count=0, + unprocessed_count=0, + total_count=5, + distributed_map_run_arn="arn:aws:lambda:us-east-1:123456789012:function:fn:$LATEST/durable-execution/exec1/invoke1/distributed-map-run/abc", + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + result = distributed_map_handler( + source=["a", "b"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr1"), + ) + + assert result.status is DistributedMapStatus.SUCCEEDED + assert result.completion_reason is DistributedMapCompletionReason.ALL_COMPLETED + assert result.success_count == 5 + assert result.failure_count == 0 + assert result.total_count == 5 + assert result.distributed_map_run_arn.endswith("/distributed-map-run/abc") + mock_state.create_checkpoint.assert_not_called() + + +def test_map_run_handler_resolves_non_success_without_raising(): + """Test a non-SUCCEEDED run resolves with a summary rather than raising. + + The run failed, so distributed_map must still return the summary carrying + the counts and let the caller opt into raising. + """ + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + operation = Operation( + operation_id="mr2", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.FAILED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.FAILURE_TOLERANCE_EXCEEDED, + success_count=3, + failure_count=2, + unprocessed_count=1, + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + result = distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr2"), + ) + + assert result.status is DistributedMapStatus.FAILED + assert ( + result.completion_reason + is DistributedMapCompletionReason.FAILURE_TOLERANCE_EXCEEDED + ) + assert result.failure_count == 2 + assert result.has_failure is True + + +def test_map_run_handler_succeeded_no_details_raises(): + """Test a succeeded operation carrying no DistributedMapDetails raises ExecutionError.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + operation = Operation( + operation_id="mr3", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=None, + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + with pytest.raises(ExecutionError): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr3"), + ) + + +def test_map_run_handler_already_started(): + """Test distributed_map_handler suspends when the operation is already started.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + operation = Operation( + operation_id="mr5", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + with pytest.raises( + SuspendExecution, match="Map run mr5 started, suspending for completion" + ): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr5"), + ) + + +def test_map_run_handler_new_operation(): + """Test distributed_map_handler creates a START checkpoint for a new operation.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr6", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + with pytest.raises(SuspendExecution): + distributed_map_handler( + source=["a", "b", "c"], + processor="test_processor", + max_concurrency=42, + state=mock_state, + operation_identifier=_identifier("mr6"), + ) + + mock_state.create_checkpoint.assert_called_once() + operation_update = mock_state.create_checkpoint.call_args[1]["operation_update"] + assert operation_update.operation_id == "mr6" + assert operation_update.operation_type == OperationType.DISTRIBUTED_MAP + assert operation_update.action == OperationAction.START + assert operation_update.name == "test_map_run" + + distributed_map_options = operation_update.to_dict()["DistributedMapOptions"] + assert distributed_map_options["MaxConcurrency"] == 42 + assert distributed_map_options["Processor"]["FunctionName"] == "test_processor" + assert distributed_map_options["Source"]["InlineSourceConfig"]["Items"] == [ + '"a"', + '"b"', + '"c"', + ] + + +def test_map_run_handler_no_config(): + """Test distributed_map_handler uses a default config when none is provided.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr7", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + with pytest.raises(SuspendExecution): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr7"), + config=None, + ) + + mock_state.create_checkpoint.assert_called_once() + + +# Immediate Response Handling Tests +# ============================================================================ + + +def test_map_run_immediate_response_get_checkpoint_result_called_twice(): + """Test get_checkpoint_result is called twice when a checkpoint is created.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr8", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + with pytest.raises(SuspendExecution): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr8"), + ) + + assert mock_state.get_checkpoint_result.call_count == 2 + + +def test_map_run_immediate_response_create_checkpoint_is_sync_true(): + """Test create_checkpoint is called with is_sync=True.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr9", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + with pytest.raises(SuspendExecution): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr9"), + ) + + mock_state.create_checkpoint.assert_called_once() + assert mock_state.create_checkpoint.call_args[1]["is_sync"] is True + + +def test_map_run_immediate_response_immediate_success(): + """Test immediate success: second check returns SUCCEEDED, summary returned.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + succeeded_op = Operation( + operation_id="mr10", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + ), + ) + succeeded = CheckpointedResult.create_from_operation(succeeded_op) + mock_state.get_checkpoint_result.side_effect = [not_found, succeeded] + + result = distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr10"), + ) + + assert result.status is DistributedMapStatus.SUCCEEDED + assert result.success_count == 1 + mock_state.create_checkpoint.assert_called_once() + assert mock_state.get_checkpoint_result.call_count == 2 + + +def test_map_run_immediate_response_no_immediate_response(): + """Test no immediate response: second check returns STARTED, suspends.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr12", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + with pytest.raises(SuspendExecution): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr12"), + ) + + mock_state.create_checkpoint.assert_called_once() + assert mock_state.get_checkpoint_result.call_count == 2 + + +def test_map_run_immediate_response_already_completed(): + """Test already completed: first check is SUCCEEDED, no checkpoint created.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + succeeded_op = Operation( + operation_id="mr13", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(succeeded_op) + ) + + result = distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr13"), + ) + + assert result.status is DistributedMapStatus.SUCCEEDED + mock_state.create_checkpoint.assert_not_called() + assert mock_state.get_checkpoint_result.call_count == 1 + + +@patch( + "aws_durable_execution_sdk_python.operation.dmap.suspend_with_optional_resume_delay" +) +def test_map_run_handler_suspend_does_not_raise(mock_suspend): + """Test distributed_map_handler raises ExecutionError if suspend does not raise.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + + not_found = CheckpointedResult.create_not_found() + started_op = Operation( + operation_id="mr14", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + started = CheckpointedResult.create_from_operation(started_op) + mock_state.get_checkpoint_result.side_effect = [not_found, started] + + mock_suspend.return_value = None + + with pytest.raises( + ExecutionError, + match="suspend_with_optional_resume_delay should have raised an exception, but did not.", + ): + distributed_map_handler( + source=["a"], + processor="test_processor", + max_concurrency=10, + state=mock_state, + operation_identifier=_identifier("mr14"), + ) + + mock_suspend.assert_called_once() + + +# Serialization and result-collection tests +# ============================================================================ + + +def _start_options(state_calls, source, processor, max_concurrency, config): + """Run the executor through the new-operation path and return the sent options dict.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + mock_state.get_checkpoint_result.side_effect = state_calls + + executor = DistributedMapOperationExecutor( + source=source, + processor=processor, + max_concurrency=max_concurrency, + state=mock_state, + operation_identifier=_identifier("mrw"), + config=config, + ) + with pytest.raises(SuspendExecution): + executor.process() + update = mock_state.create_checkpoint.call_args[1]["operation_update"] + return update.to_dict()["DistributedMapOptions"] + + +def _new_op_state_calls(): + not_found = CheckpointedResult.create_not_found() + started = CheckpointedResult.create_from_operation( + Operation( + operation_id="mrw", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STARTED, + ) + ) + return [not_found, started] + + +def test_processor_item_failures_sets_response_types(): + """item_failures serializes FunctionResponseTypes=REPORT_BATCH_ITEM_FAILURES.""" + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.item_failures("proc", batch_size=25), + max_concurrency=4, + config=DistributedMapConfig(), + ) + processor = options["Processor"] + assert processor["FunctionName"] == "proc" + assert processor["FunctionResponseTypes"] == ["REPORT_BATCH_ITEM_FAILURES"] + assert processor["BatchSize"] == 25 + + +def test_processor_unlimited_retries_maps_to_negative_one(): + """DistributedMapProcessor.UNLIMITED serializes to MaxRetryAttempts=-1.""" + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.item_results( + "proc", + max_retry_attempts=DistributedMapProcessor.UNLIMITED, + max_retry_duration=Duration.from_hours(1), + ), + max_concurrency=1, + config=DistributedMapConfig(), + ) + processor = options["Processor"] + assert processor["FunctionResponseTypes"] == ["REPORT_BATCH_ITEM_RESULTS"] + assert processor["MaxRetryAttempts"] == -1 + assert processor["MaxRetryDurationSeconds"] == 3600 + + +def test_processor_explicit_retry_attempts_pass_through(): + """A plain int retry count passes through unchanged (no sentinel mapping).""" + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("proc", max_retry_attempts=0), + max_concurrency=1, + config=DistributedMapConfig(), + ) + assert options["Processor"]["MaxRetryAttempts"] == 0 + # batch mode reports no per-item response types + assert "FunctionResponseTypes" not in options["Processor"] + + +def test_s3_source_serializes_config(): + """An S3 json_lines source serializes to an S3SourceConfig block.""" + options = _start_options( + _new_op_state_calls(), + source=S3Source.json_lines("s3://bucket/data.jsonl", max_items=500), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=2, + config=DistributedMapConfig(), + ) + source = options["Source"] + assert source["Type"] == "S3" + assert source["MaxItemsToRead"] == 500 + assert source["S3SourceConfig"]["Bucket"] == "bucket" + assert source["S3SourceConfig"]["Key"] == "data.jsonl" + assert source["S3SourceConfig"]["Format"] == "JSON_LINES" + + +def test_full_config_serializes_all_blocks(): + """Completion, destination, timeout, and result-collection blocks all serialize.""" + config = DistributedMapResultConfig( + completion_config=DistributedMapCompletionConfig.failure_percentage( + 5, minimum_sample_size=200 + ), + destination=DistributedMapDestinationConfig( + on_success=S3Destination.successes("s3://out/ok"), + on_failure=S3Destination.failures("s3://out/bad"), + ), + timeout=Duration.from_minutes(30), + ) + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.item_results("proc"), + max_concurrency=8, + config=config, + ) + assert options["CompletionConfig"] == { + "ToleratedFailurePercentage": 5, + "MinimumSampleSize": 200, + } + on_success = options["Destination"]["OnSuccess"] + assert on_success["Type"] == "S3" + assert on_success["S3DestinationConfig"]["Bucket"] == "out" + assert on_success["S3DestinationConfig"]["KeyPrefix"] == "ok" + on_failure = options["Destination"]["OnFailure"] + assert on_failure["Include"] == ["INPUT", "ERROR"] + assert on_failure["S3DestinationConfig"]["Bucket"] == "out" + assert options["TimeoutSeconds"] == 1800 + assert options["ResultCollection"] == {"Mode": "INLINE"} + + +def test_result_config_returns_map_run_result_with_items(): + """With a DistributedMapResultConfig, a DistributedMapResult with per-item results is built.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + error = ErrorObject(message="boom", type="ItemError", data=None, stack_trace=None) + operation = Operation( + operation_id="mrr", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=1, + unprocessed_count=0, + total_count=2, + results=( + DistributedMapResultItemApi( + item_id="0", status=DistributedMapItemStatus.SUCCEEDED, output="42" + ), + DistributedMapResultItemApi( + item_id="1", status=DistributedMapItemStatus.FAILED, error=error + ), + ), + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + result = distributed_map_handler( + source=["a", "b"], + processor=DistributedMapProcessor.item_results("proc"), + max_concurrency=2, + state=mock_state, + operation_identifier=_identifier("mrr"), + config=DistributedMapResultConfig(), + ) + + assert isinstance(result, DistributedMapResult) + assert len(result.all) == 2 + assert result.get_results() == [42] + errors = result.get_errors() + assert len(errors) == 1 + assert errors[0].error_type == "ItemError" + assert errors[0].error_message == "boom" + + +def test_plain_config_returns_plain_summary(): + """With a plain DistributedMapConfig, a DistributedMapSummary is returned.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="mrs", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=2, + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + + result = distributed_map_handler( + source=["a", "b"], + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=2, + state=mock_state, + operation_identifier=_identifier("mrs"), + ) + + assert isinstance(result, DistributedMapSummary) + assert not isinstance(result, DistributedMapResult) + assert result.success_count == 2 + + +def test_csv_source_header_location(): + """CSV headers map to GIVEN. expected_columns stays client-side under FIRST_ROW.""" + # headers -> HeaderLocation GIVEN, headers sent + given = _start_options( + _new_op_state_calls(), + source=S3Source.csv("s3://b/data.csv", headers=["a", "b"]), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + csv_opts = given["Source"]["S3SourceConfig"]["CsvFormatOptions"] + assert csv_opts["HeaderLocation"] == "GIVEN" + assert csv_opts["Headers"] == ["a", "b"] + assert csv_opts["Delimiter"] == "COMMA" + + # no headers -> HeaderLocation FIRST_ROW + first_row = _start_options( + _new_op_state_calls(), + source=S3Source.csv("s3://b/data.csv"), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + s3_cfg = first_row["Source"]["S3SourceConfig"] + assert s3_cfg["CsvFormatOptions"]["HeaderLocation"] == "FIRST_ROW" + assert "Headers" not in s3_cfg["CsvFormatOptions"] + + +# Call-site validation and serdes tests +# ============================================================================ + + +def test_csv_delimiter_accepts_enum(): + opts = _start_options( + _new_op_state_calls(), + source=S3Source.csv( + "s3://b/data.csv", delimiter=DistributedMapCsvDelimiter.PIPE + ), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + assert opts["Source"]["S3SourceConfig"]["CsvFormatOptions"]["Delimiter"] == "PIPE" + + +def test_inline_source_over_1mb_rejected(): + big = ["x" * 100_000] * 12 # ~1.2 MB serialized + with pytest.raises(ValidationError, match="1 MB limit"): + _start_options( + _new_op_state_calls(), + source=big, + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + + +def test_reader_state_serialized_and_capped(): + # typed initial_state is serialized into the opaque state string + options = _start_options( + _new_op_state_calls(), + source=ReaderSource.from_function("reader", initial_state={"page": 0}), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + reader_cfg = options["Source"]["ReaderFunctionSourceConfig"] + assert reader_cfg["FunctionName"] == "reader" + assert reader_cfg["InitialState"] == '{"page": 0}' + + # oversize state is rejected + with pytest.raises(ValidationError, match="32 KB limit"): + _start_options( + _new_op_state_calls(), + source=ReaderSource.from_function("reader", initial_state="x" * 40_000), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + + +def test_inline_custom_serdes_applied_to_items(): + """A custom inline serdes transforms each item's serialized value.""" + from aws_durable_execution_sdk_python.serdes import SerDes + + class _UpperSerDes(SerDes): + def serialize(self, value, _serdes_context): # noqa: ANN001, ANN201 + return json.dumps(value.upper()) + + def deserialize(self, data, _serdes_context): # noqa: ANN001, ANN201 + return json.loads(data).lower() + + options = _start_options( + _new_op_state_calls(), + source=InlineSource.of(["a", "b"], serdes=_UpperSerDes()), + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + assert options["Source"]["InlineSourceConfig"]["Items"] == ['"A"', '"B"'] + + +# Result helpers, from_dict, destinations, and source variants +# ============================================================================ + + +def test_map_run_options_from_dict_round_trip(): + sent = _start_options( + _new_op_state_calls(), + source=["a", "b"], + processor=DistributedMapProcessor.batch("proc"), + max_concurrency=7, + config=DistributedMapConfig(), + ) + parsed = DistributedMapOptions.from_dict(sent) + assert parsed.max_concurrency == 7 + assert parsed.source.source_type is DistributedMapSourceType.INLINE + assert parsed.processor.function_name == "proc" + + +def test_destination_only_success(): + opts = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_success=S3Destination.successes("s3://out/ok") + ) + ), + ) + dest = opts["Destination"] + assert "OnSuccess" in dest + assert "OnFailure" not in dest + + +def test_destination_only_failure(): + opts = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_failure=S3Destination.failures("s3://out/bad") + ) + ), + ) + dest = opts["Destination"] + assert "OnFailure" in dest + assert "OnSuccess" not in dest + + +def test_s3_objects_source_config(): + opts = _start_options( + _new_op_state_calls(), + source=S3Source.objects("s3://b/prefix/"), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + s3 = opts["Source"]["S3SourceConfig"] + assert s3["Transform"] == "NONE" + assert s3["KeyPrefix"] == "prefix/" + assert "Format" not in s3 + + +def test_s3_flattened_json_lines_source_config(): + opts = _start_options( + _new_op_state_calls(), + source=S3Source.flattened_json_lines("s3://b/prefix/"), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + s3 = opts["Source"]["S3SourceConfig"] + assert s3["Transform"] == "LOAD_AND_FLATTEN" + assert s3["Format"] == "JSON_LINES" + + +# Validations, config branches, and round-trips +# ============================================================================ + + +def _start_executor(source, processor, config, max_concurrency=1): + """Build an executor on the new-operation path (for error-path tests).""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + mock_state.get_checkpoint_result.side_effect = _new_op_state_calls() + return DistributedMapOperationExecutor( + source=source, + processor=processor, + max_concurrency=max_concurrency, + state=mock_state, + operation_identifier=_identifier("mrw"), + config=config, + ) + + +# --- Completion config translation --- + + +def test_empty_completion_config_omitted(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(completion_config=DistributedMapCompletionConfig()), + ) + assert "CompletionConfig" not in options + + +# --- Source translation variants --- + + +def test_objects_whole_bucket_prefix_config(): + options = _start_options( + _new_op_state_calls(), + source=S3Source.objects("s3://bucket"), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + s3 = options["Source"]["S3SourceConfig"] + assert s3["KeyPrefix"] == "" + assert s3["Transform"] == "NONE" + assert "Key" not in s3 + + +def test_flattened_csv_source_config(): + options = _start_options( + _new_op_state_calls(), + source=S3Source.flattened_csv("s3://b/prefix", headers=["a", "b"]), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + s3 = options["Source"]["S3SourceConfig"] + assert s3["Transform"] == "LOAD_AND_FLATTEN" + assert s3["Format"] == "CSV" + assert s3["CsvFormatOptions"]["HeaderLocation"] == "GIVEN" + assert s3["CsvFormatOptions"]["Headers"] == ["a", "b"] + + +def test_reader_source_without_initial_state_omits_state(): + options = _start_options( + _new_op_state_calls(), + source=ReaderSource.from_function("reader"), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + reader = options["Source"]["ReaderFunctionSourceConfig"] + assert reader["FunctionName"] == "reader" + assert "InitialState" not in reader + + +def test_inline_non_json_serdes_supported(): + """Items are opaque strings, so a serdes need not produce JSON.""" + from aws_durable_execution_sdk_python.serdes import SerDes + + class _PlainTextSerDes(SerDes): + def serialize(self, value, _serdes_context): # noqa: ANN001, ANN201 + return f"item-{value}" + + def deserialize(self, data, _serdes_context): # noqa: ANN001, ANN201, ARG002 + return data + + options = _start_options( + _new_op_state_calls(), + source=InlineSource.of([1, 2], serdes=_PlainTextSerDes()), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + assert options["Source"]["InlineSourceConfig"]["Items"] == ["item-1", "item-2"] + + +# --- Destination translation permutations --- + + +def test_success_destination_include_input_and_owner(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_success=S3Destination.successes( + "s3://out/ok", + include_input=True, + include_output=True, + expected_bucket_owner="123456789012", + ) + ) + ), + ) + on_success = options["Destination"]["OnSuccess"] + assert on_success["Include"] == ["INPUT", "OUTPUT"] + assert on_success["S3DestinationConfig"]["ExpectedBucketOwner"] == "123456789012" + assert "OnFailure" not in options["Destination"] + + +def test_failure_destination_error_only_and_owner(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_failure=S3Destination.failures( + "s3://out/bad", + include_input=False, + include_error=True, + expected_bucket_owner="123456789012", + ) + ) + ), + ) + on_failure = options["Destination"]["OnFailure"] + assert on_failure["Include"] == ["ERROR"] + assert on_failure["S3DestinationConfig"]["ExpectedBucketOwner"] == "123456789012" + + +# --- Operation round-trips through the checkpoint --- + + +def test_operation_to_dict_round_trip_preserves_results(): + op = Operation( + operation_id="opx", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=1, + unprocessed_count=0, + total_count=2, + distributed_map_run_arn="arn:aws:lambda:us-east-1:123456789012:function:fn:$LATEST/durable-execution/exec1/invoke1/distributed-map-run/z", + completion_details="done", + results=( + DistributedMapResultItemApi( + item_id="0", status=DistributedMapItemStatus.SUCCEEDED, output=5 + ), + DistributedMapResultItemApi( + item_id="1", + status=DistributedMapItemStatus.FAILED, + error=ErrorObject( + message="boom", type="E", data=None, stack_trace=None + ), + ), + ), + ), + ) + block = op.to_dict()["DistributedMapDetails"] + assert block["Results"][0] == {"ItemId": "0", "Status": "SUCCEEDED", "Output": 5} + assert block["DistributedMapRunArn"].endswith("/distributed-map-run/z") + assert block["CompletionDetails"] == "done" + assert block["TotalCount"] == 2 + + parsed = Operation.from_dict(op.to_dict()) + assert parsed.distributed_map_details is not None + assert parsed.distributed_map_details.results[0].output == 5 + assert parsed.distributed_map_details.results[1].error.type == "E" + + +def test_options_full_round_trip(): + options_dict = _start_options( + _new_op_state_calls(), + source=S3Source.csv("s3://b/f.csv", headers=["a"]), + processor=DistributedMapProcessor.item_results( + "proc", + batch_size=5, + max_retry_attempts=DistributedMapProcessor.UNLIMITED, + max_retry_duration=Duration.from_minutes(10), + durable_execution_name_prefix="pfx", + ), + max_concurrency=3, + config=DistributedMapResultConfig( + destination=DistributedMapDestinationConfig( + on_success=S3Destination.successes("s3://o/ok"), + on_failure=S3Destination.failures("s3://o/bad"), + ), + completion_config=DistributedMapCompletionConfig.failure_count(2), + timeout=Duration.from_minutes(5), + ), + ) + parsed = DistributedMapOptions.from_dict(options_dict) + assert parsed.max_concurrency == 3 + assert parsed.source.source_type is DistributedMapSourceType.S3 + assert parsed.processor.function_response_types == ( + DistributedMapFunctionResponseType.REPORT_BATCH_ITEM_RESULTS, + ) + assert parsed.processor.max_retry_attempts == -1 + assert parsed.processor.durable_execution_name_prefix == "pfx" + assert parsed.destination is not None + assert parsed.completion_config.tolerated_failure_count == 2 + assert parsed.result_collection.mode is DistributedMapResultCollectionMode.INLINE + assert parsed.timeout_seconds == 300 + assert parsed.to_dict()["MaxConcurrency"] == 3 + + +# Remaining validation and config branches +# ============================================================================ + + +def test_csv_first_row_no_headers(): + options = _start_options( + _new_op_state_calls(), + source=S3Source.csv("s3://b/f.csv"), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + csv_opts = options["Source"]["S3SourceConfig"]["CsvFormatOptions"] + assert csv_opts["HeaderLocation"] == "FIRST_ROW" + assert "Headers" not in csv_opts + assert csv_opts["Delimiter"] == "COMMA" + + +def test_s3_source_expected_bucket_owner(): + options = _start_options( + _new_op_state_calls(), + source=S3Source.json_lines( + "s3://b/k.jsonl", expected_bucket_owner="123456789012" + ), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + assert options["Source"]["S3SourceConfig"]["ExpectedBucketOwner"] == "123456789012" + + +def test_empty_destination_config_omitted(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(destination=DistributedMapDestinationConfig()), + ) + assert "Destination" not in options + + +def test_success_destination_input_only(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_success=S3Destination.successes( + "s3://out/ok", include_input=True, include_output=False + ) + ) + ), + ) + assert options["Destination"]["OnSuccess"]["Include"] == ["INPUT"] + + +def test_failure_destination_input_only(): + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig( + destination=DistributedMapDestinationConfig( + on_failure=S3Destination.failures( + "s3://out/bad", include_input=True, include_error=False + ) + ) + ), + ) + assert options["Destination"]["OnFailure"]["Include"] == ["INPUT"] + + +def test_reader_source_options_round_trip(): + options_dict = _start_options( + _new_op_state_calls(), + source=ReaderSource.from_function("reader", initial_state={"page": 0}), + processor=DistributedMapProcessor.batch("p"), + max_concurrency=1, + config=DistributedMapConfig(), + ) + parsed = DistributedMapOptions.from_dict(options_dict) + assert parsed.source.source_type is DistributedMapSourceType.READER_FUNCTION + assert parsed.source.reader_config is not None + assert parsed.source.reader_config.function_name == "reader" + + +def test_operation_to_dict_minimal_details_omits_optionals(): + op = Operation( + operation_id="opm", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=0, + unprocessed_count=0, + ), + ) + block = op.to_dict()["DistributedMapDetails"] + assert "DistributedMapRunArn" not in block + assert "CompletionDetails" not in block + assert "TotalCount" not in block + assert "Results" not in block + + +@pytest.mark.parametrize( + "status", + [OperationStatus.FAILED, OperationStatus.STOPPED, OperationStatus.TIMED_OUT], +) +def test_operation_level_terminal_failure_without_details_raises(status): + """A terminal failure carrying no details raises, since counts cannot be built.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="mrf", + operation_type=OperationType.DISTRIBUTED_MAP, + status=status, + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + with pytest.raises(ExecutionError): + distributed_map_handler( + source=["a"], + processor="p", + max_concurrency=1, + state=mock_state, + operation_identifier=_identifier("mrf"), + ) + + +@pytest.mark.parametrize( + ("status", "expected"), + [ + (OperationStatus.FAILED, DistributedMapStatus.FAILED), + (OperationStatus.TIMED_OUT, DistributedMapStatus.TIMED_OUT), + (OperationStatus.STOPPED, DistributedMapStatus.STOPPED), + ], +) +def test_terminal_failure_with_details_resolves_with_summary(status, expected): + """A terminal non-success run resolves with its summary rather than raising.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="mrf", + operation_type=OperationType.DISTRIBUTED_MAP, + status=status, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.FAILURE_TOLERANCE_EXCEEDED, + success_count=1, + failure_count=2, + unprocessed_count=0, + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + summary = distributed_map_handler( + source=["a"], + processor="p", + max_concurrency=1, + state=mock_state, + operation_identifier=_identifier("mrf"), + ) + assert summary.status is expected + assert summary.failure_count == 2 + # The caller opts into raising rather than having it forced on them. + with pytest.raises(DistributedMapError): + summary.throw_if_error() + + +def test_terminal_without_completion_reason_raises(): + """A terminal operation missing its completion reason is surfaced, not papered over.""" + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="mrf", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.STOPPED, + distributed_map_details=DistributedMapDetails( + success_count=0, + failure_count=0, + unprocessed_count=3, + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + with pytest.raises(ExecutionError, match="carried no CompletionReason"): + distributed_map_handler( + source=["a"], + processor="p", + max_concurrency=1, + state=mock_state, + operation_identifier=_identifier("mrf"), + ) + + +# Serdes decode, falsy output, and duration-only retry +# ============================================================================ + + +def test_custom_result_serdes_applied_on_decode(): + """A custom result_serdes transforms each per-item output on decode.""" + from aws_durable_execution_sdk_python.serdes import SerDes + + class _UpperSerDes(SerDes): + def serialize(self, value, _serdes_context): # noqa: ANN001, ANN201 + return json.dumps(value) + + def deserialize(self, data, _serdes_context): # noqa: ANN001, ANN201 + return json.loads(data).upper() + + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="cs", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=0, + unprocessed_count=0, + results=( + DistributedMapResultItemApi( + item_id="0", + status=DistributedMapItemStatus.SUCCEEDED, + output='"abc"', + ), + ), + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + result = distributed_map_handler( + source=["a"], + processor=DistributedMapProcessor.item_results("proc"), + max_concurrency=1, + state=mock_state, + operation_identifier=_identifier("cs"), + config=DistributedMapResultConfig(result_serdes=_UpperSerDes()), + ) + assert result.get_results() == ["ABC"] + + +@pytest.mark.parametrize("value", [0, False, "", [], {}]) +def test_falsy_output_round_trips(value): + """A falsy-but-present output survives the round-trip and decode (not dropped).""" + entry = DistributedMapResultItemApi( + item_id="0", + status=DistributedMapItemStatus.SUCCEEDED, + output=json.dumps(value), + ) + assert entry.to_dict()["Output"] == json.dumps(value) + + mock_state = Mock(spec=ExecutionState) + mock_state.durable_execution_arn = "test_arn" + operation = Operation( + operation_id="fo", + operation_type=OperationType.DISTRIBUTED_MAP, + status=OperationStatus.SUCCEEDED, + distributed_map_details=DistributedMapDetails( + completion_reason=DistributedMapCompletionReason.ALL_COMPLETED, + success_count=1, + failure_count=0, + unprocessed_count=0, + results=(entry,), + ), + ) + mock_state.get_checkpoint_result.return_value = ( + CheckpointedResult.create_from_operation(operation) + ) + result = distributed_map_handler( + source=["a"], + processor=DistributedMapProcessor.item_results("proc"), + max_concurrency=1, + state=mock_state, + operation_identifier=_identifier("fo"), + config=DistributedMapResultConfig(), + ) + assert result.get_results() == [value] + + +def test_processor_retry_duration_only(): + """A retry config with only a duration sends the duration and omits attempts.""" + options = _start_options( + _new_op_state_calls(), + source=["a"], + processor=DistributedMapProcessor.batch( + "proc", + max_retry_duration=Duration.from_minutes(5), + ), + max_concurrency=1, + config=DistributedMapConfig(), + ) + processor = options["Processor"] + assert processor["MaxRetryDurationSeconds"] == 300 + assert "MaxRetryAttempts" not in processor