diff --git a/.github/workflows/lmi-e2e-tests.yml b/.github/workflows/lmi-e2e-tests.yml new file mode 100644 index 000000000..84852a210 --- /dev/null +++ b/.github/workflows/lmi-e2e-tests.yml @@ -0,0 +1,101 @@ +name: LMI E2E Tests + +on: + push: + branches: [main] + workflow_dispatch: + pull_request: + types: [opened, synchronize, reopened, labeled] + paths: + - 'lmi-tests/**' + - 'sdk/src/test/java/software/amazon/lambda/durable/execution/LmiLifecycleRegressionTest.java' + - '.github/workflows/lmi-e2e-tests.yml' + +permissions: + contents: read + id-token: write + +# Persistent functions are shared; serialize deployment and testing across all runs. +concurrency: + group: java-lmi-e2e + cancel-in-progress: false + +jobs: + cloud: + if: >- + (github.event_name != 'pull_request' || + (github.event.pull_request.head.repo.full_name == github.repository && + github.actor != 'dependabot[bot]' && + contains(github.event.pull_request.labels.*.name, 'run-lmi-e2e'))) + runs-on: ubuntu-latest + timeout-minutes: 85 + env: + LMI_STACK_NAME: java-lmi-e2e + CAPACITY_PROVIDER_ARN: ${{ secrets.CAPACITY_PROVIDER_ARN }} + TEST_LAMBDA_EXECUTION_ROLE_ARN: ${{ secrets.TEST_LAMBDA_EXECUTION_ROLE_ARN }} + steps: + - uses: actions/checkout@v7 + - name: Validate configuration + shell: bash + run: | + set -euo pipefail + test -n "$CAPACITY_PROVIDER_ARN" + test -n "$TEST_LAMBDA_EXECUTION_ROLE_ARN" + region=$(cut -d: -f4 <<< "$CAPACITY_PROVIDER_ARN") + echo "AWS_REGION=$region" >> "$GITHUB_ENV" + - uses: aws-actions/configure-aws-credentials@e1253824e5c10ff9df46874f81ed3ec929e19cfd # v6.3.0 + with: + role-to-assume: ${{ secrets.TEST_ROLE_ARN }} + role-session-name: java-lmi-e2e + aws-region: ${{ env.AWS_REGION }} + allowed-account-ids: ${{ secrets.TEST_ACCOUNT_ID }} + - uses: actions/setup-java@v6 + with: + distribution: corretto + java-version: '25' + cache: maven + - name: Build commit under test and validate evidence assertions + timeout-minutes: 10 + run: | + mvn -B -q -pl lmi-tests -am package + mvn -B -q -pl lmi-tests -am spotless:check + python3 -m unittest discover -s lmi-tests/tests -v + - name: Create or update persistent LMI functions + timeout-minutes: 35 + run: python3 lmi-tests/cloud_suite.py deploy --stack-name "$LMI_STACK_NAME" --run-id "${GITHUB_RUN_ID}-${GITHUB_RUN_ATTEMPT}" + - name: LMI cloud regression assertions (expected red until issue 726 is fixed) + timeout-minutes: 30 + run: python3 -u lmi-tests/cloud_suite.py test --cloud-enabled + - name: Collect histories and lifecycle evidence + if: always() + timeout-minutes: 5 + run: | + if test -f lmi-tests/artifacts/manifest.json; then + python3 lmi-tests/cloud_suite.py collect + fi + - name: Publish failure evidence + if: always() + uses: actions/upload-artifact@v7 + with: + name: lmi-e2e-${{ github.run_id }}-${{ github.run_attempt }} + retention-days: 7 + path: | + lmi-tests/artifacts/ + lmi-tests/target/surefire-reports/ + sdk/target/surefire-reports/ + - name: Summarize outcomes + if: always() + run: | + python3 - <<'PY' + import os, pathlib, xml.etree.ElementTree as ET + p = pathlib.Path('lmi-tests/artifacts/junit.xml') + with open(os.environ['GITHUB_STEP_SUMMARY'], 'a') as out: + out.write('LMI tests assert the fixed behavior in #726; failures are not expected-success results.\n\n') + if p.exists(): + for case in ET.parse(p).getroot().findall('testcase'): + failures = list(case) + state = 'PASS' if not failures else failures[0].get('type', 'FAIL') + out.write(f"- {case.get('name')}: {state}\n") + else: + out.write('Cloud assertions did not run. Inspect setup errors in the artifacts.\n') + PY diff --git a/lmi-tests/.gitignore b/lmi-tests/.gitignore new file mode 100644 index 000000000..2a2d39e00 --- /dev/null +++ b/lmi-tests/.gitignore @@ -0,0 +1,3 @@ +/artifacts/ +/__pycache__/ +/tests/__pycache__/ diff --git a/lmi-tests/README.md b/lmi-tests/README.md new file mode 100644 index 000000000..75ad3bfcd --- /dev/null +++ b/lmi-tests/README.md @@ -0,0 +1,178 @@ +# LMI lifecycle cloud tests + +This suite tests the SDK commit being built on real Lambda Managed +Instances (LMI). It asserts the desired behavior in [#726](https://github.com/aws/aws-durable-execution-sdk-java/issues/726) +and implements the cloud coverage requested in [#727](https://github.com/aws/aws-durable-execution-sdk-java/issues/727). +The affected SDK is expected to fail. Do not invert assertions, skip regressions, +or accept a retry that happens to pass after a lifecycle violation. + +## Test design + +* Each fixture is a published durable function with LMI invocation concurrency + 1, 2, or 8. Java 25 / arm64 is the initial matrix. Java 17 is not supported by + LMI. Deployment and readback are the region/architecture capability check: + unsupported combinations fail setup; there is no ordinary-Lambda fallback. +* One CI job builds once, deploys all five fixture functions, runs all 13 cases, + then collects evidence. A persistent CloudFormation stack owns all five + functions and log groups. The stack, functions and bucket remain after both + successful and failed test runs; later deployments update the same resources. Each uses 2 GiB / 1 vCPU and a single + `$LATEST.PUBLISHED` version; code is not republished during the test phase, and + the digest is verified against the built artifact. Deployment and tests are + serialized. Before updating an existing stack, the driver waits for earlier + durable executions on its fixture functions to finish. Each deployment supplies + `LMI_TEST_RUN_ID` to refresh fault-fixture JVM state (including an executor grown + by a prior escape), while keeping resource names and ARNs stable. Runtime traces + must match the current deployment run ID and commit. +* The template declares `FunctionScalingConfig` with both + `MinExecutionEnvironments` and `MaxExecutionEnvironments` set to 1 on every + function. CloudFormation applies those limits as part of resource creation, + so setup does not first stabilize functions with the default three-environment + floor and then lower it. The driver reads back the applied limits and ACTIVE + version state before testing. The provider's 12-vCPU limit applies to EC2 + instance capacity, including placement and instance overhead; the suite never + changes that limit. Invocation concurrency remains 1, 2, or 8 per environment, + and same-JVM overlap must still be proven. +* A stream wrapper observes the actual SDK entry and return. Invocation-local + root/task `finally` markers and a JVM-wide sequence establish ordering. The + plugin end hook is deliberately not used as a completion signal. +* A bounded same-JVM root barrier establishes the fixed-pool reproduction. + Other cases hold steps with private S3 control objects. The driver requires + distinct request IDs active in the same JVM. Placement has its own deadline + and failure category. Environment replacement is never worker recovery. +* Timeout victims start before healthy peers, leaving the peers time to hold + the other runtime slots during recovery. Diagnostics capture the real context + deadline; no fake clock, context, checkpoint backend, or time-skipping runner + participates. Service timeout evidence, task interruption/exit, wrapper exit, + and restored admission are independent assertions. Durable execution status + is collected separately from the runtime invocation's outcome. +* Successful and failed checkpointed steps precede a real durable wait. The + attempt ledger records entry into user bodies, while real service history + proves checkpoint identity and replay. Fixture business work returns an idempotent marker; BODY events form the + append-only attempt ledger keyed by execution and operation. Interrupted, uncheckpointed work may be retried. +* Diagnostics (including target-JVM admission and barriers) are test-only + instrumentation. They never choose operation names or business branches. + Deliberately blocked tasks have finite escape timers. Escape diagnostics fail + lifecycle assertions; they cannot turn the reproduced bug into a pass. +* CloudWatch collection polls for causal evidence and deduplicates JVM sequence + numbers. Test reports distinguish setup, placement, assertion, collection, + and teardown failures. Raw histories, configuration, and diagnostic logs are + retained even when a scenario fails. + +## Invocation qualifier + +LMI requires a published version. The suite uses the mutable `$LATEST.PUBLISHED` +qualifier to update the same test functions without accumulating numbered +versions. The unpublished `$LATEST` qualifier allowed by general Durable +Functions documentation is not an LMI deployment target. An unqualified LMI ARN +also resolves to `$LATEST.PUBLISHED`. See the LMI-specific +[getting-started requirement](https://docs.aws.amazon.com/lambda/latest/dg/lambda-managed-instances-getting-started.html) +and [version behavior](https://docs.aws.amazon.com/lambda/latest/dg/lambda-managed-instances-version-publishing.html). + +## Ownership + +`CAPACITY_PROVIDER_ARN` identifies an existing **dedicated test** capacity +provider. The suite never creates, updates, or deletes it. Its owner must bound +its maximum vCPUs and provide working Lambda/S3/CloudWatch connectivity. The +workflow uses `TEST_ROLE_ARN`, `TEST_ACCOUNT_ID`, and +`TEST_LAMBDA_EXECUTION_ROLE_ARN`, as the ordinary E2E workflow does. + +The default persistent stack is `java-lmi-e2e`, with functions +`java-lmi-e2e-default1`, `java-lmi-e2e-default2`, `java-lmi-e2e-default8`, +`java-lmi-e2e-fixed2`, and `java-lmi-e2e-nested2`. Use `--stack-name` to select +another dedicated stack. The private bucket name is stable for the account, +region and stack. Existing resources must carry this suite's ownership tags. + +There is no automatic stack/function/bucket deletion and no janitor. The +CloudFormation service's normal deployment rollback behavior still applies. +Failed deployments remain available for diagnosis; the driver never deletes and +recreates a failed stack to conceal the failure. The test-account owner manages +retirement of the persistent resources. + +Each run uses its own `control/RUN_ID/` object prefix. Only control objects expire +after one day. SDK jars use content-addressed `code/` keys and remain available +for later deployments and CloudFormation rollback. CloudWatch logs and durable +histories retain one day; GitHub artifacts retain seven days. + +## Running + +See the workflow `lmi-e2e-tests.yml` for the complete commands and budgets. The +cloud driver requires Python 3.9+, AWS CLI v2 with LMI/Durable API support, and +credentials for the dedicated test account. There are no new Python packages. +The Java fixture uses the repository SDK and existing dependencies only. One local +contract test runs the real AWS CLI against an unsigned localhost HTTP endpoint +to validate synchronous/asynchronous invocation arguments and payload bytes. It +does not call AWS or provide cloud coverage. + +```sh +mvn -B -pl lmi-tests -am package -DskipTests +python3 -m unittest discover -s lmi-tests/tests -v +export CAPACITY_PROVIDER_ARN=arn:aws:lambda:REGION:ACCOUNT:capacity-provider:NAME +export TEST_LAMBDA_EXECUTION_ROLE_ARN=arn:aws:iam::ACCOUNT:role/ROLE +export AWS_REGION=us-west-2 +python3 lmi-tests/cloud_suite.py deploy --stack-name java-lmi-e2e --run-id local-unique +python3 lmi-tests/cloud_suite.py test --cloud-enabled +python3 lmi-tests/cloud_suite.py collect +``` + +The deploy command creates or updates and retains `default1`, `default2`, `default8`, +`fixed2`, and `nested2` together. The test command requires all five and runs the +full suite. CI publishes one combined artifact, named `lmi-e2e-RUN-ATTEMPT`, with +the single `template.json` and per-function configuration snapshots. The +CloudFormation execution role needs permission to manage the function scaling +configuration; the driver calls `lambda:GetFunctionScalingConfig` to verify it +and `lambda:ListDurableExecutionsByFunction` to wait for earlier runs before an +update. A changed jar gets a new S3 key so CloudFormation actually updates code, +not just metadata. Updates with no changes are accepted. + +Deployment records `lmi-tests/artifacts/manifest.json`, including the stack, functions, commit, jar +digest, qualified function ARNs, runtime, architecture, concurrency and provider +association. Never publish control URLs: they are temporary credentials. The +artifact writer redacts them from histories and logs. + +Cloud tests are disabled unless `test --cloud-enabled` is explicitly requested. +Local assertion tests verify that missing evidence, mismatched environments, +early responses, late tasks and stalled executors cannot be reported as passes. +Cloud regressions run on every push to `main`, including every merged change, +with no changed-path filters. Manual dispatch and same-repository PR opt-in +are also supported. There are no scheduled jobs. + +The opt-in local regressions assert the same three contracts against the SDK's +mock backend (they do not substitute for cloud coverage): + +```sh +mvn -pl sdk test -Dtest=LmiLifecycleRegressionTest -Dtest.lmi.regressions.enabled=true +``` + +For a same-repository PR, add the `run-lmi-e2e` label to opt into cloud execution. +The workflow never uses a privileged `pull_request_target` checkout. Provisioning +has a 35-minute step budget, with a shared 30-minute deadline for creating the +stack and verifying all five functions. Each scaling wait is capped at 5 minutes and at the +remaining deployment budget. Scenarios have 30 minutes, final collection 5 minutes, +with no infrastructure teardown step. Individual admission attempts are bounded (four batches, +25 seconds), fixed-pool progress has 8 seconds, and task escape timers are capped +at 120 seconds. The normal invocation timeout is 60 seconds; the durable execution +timeout is 240 seconds. Cleanup is required by the invocation deadline plus +5 seconds. Probe admission has an 8-second tolerance. Collection latency does +not extend these assertions, which compare timestamps captured inside the JVM. + +`timeouts/*.json` distinguishes server timeout logs from an SDK deadline +cancellation that returns early with an invocation error. A client HTTP timeout +is a collection error. Deadline victims use asynchronous service invocation so a +longer durable retry does not consume the driver HTTP timeout; invocation +outcomes come from request-correlated runtime logs. Invocation request failures +are saved in `invocations/*.json` with the function ARN, scenario, start time, +request state, elapsed time, and CLI exit code/error when available. The driver +checks the invocation Future during evidence polling so an API/CLI failure is +reported immediately. Missing runtime-entry evidence is a collection error; +lifecycle assertions apply after the wrapper entry has been observed. Healthy/probe scheduling +uses the actual runtime deadline, excluding cold-start and request-queue delay. +A successful durable retry cannot erase an old invocation +that exceeds the cleanup budget. No virtual-thread executor variant is deployed +until its executor contract is defined; default cached and shared fixed pools +are covered separately. + +CloudFormation's native [FunctionScalingConfig](https://docs.aws.amazon.com/AWSCloudFormation/latest/TemplateReference/aws-properties-lambda-function-functionscalingconfig.html) +sets the limits in the resource declaration. LMI provisioning and version behavior +are documented in the AWS guides for +[scaling](https://docs.aws.amazon.com/lambda/latest/dg/lambda-managed-instances-scaling.html) +and [$LATEST.PUBLISHED](https://docs.aws.amazon.com/lambda/latest/dg/lambda-managed-instances-version-publishing.html). diff --git a/lmi-tests/cloud_suite.py b/lmi-tests/cloud_suite.py new file mode 100644 index 000000000..459fa2c90 --- /dev/null +++ b/lmi-tests/cloud_suite.py @@ -0,0 +1,579 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Opt-in real-service LMI tests. See README.md for design and ownership.""" +import argparse +import base64 +import hashlib +import json +import os +from pathlib import Path +import re +import subprocess +import time +import traceback +import uuid +import xml.etree.ElementTree as ET + +from cloud_support import (Cloud, CollectionError, PreconditionError, assert_fixed, + assert_lifecycle, assert_overlap, assert_replay, aws, + require, save, selected) + +ROOT = Path(__file__).resolve().parent +ARTIFACTS = ROOT / "artifacts" +MANIFEST = ARTIFACTS / "manifest.json" +OWNER = "java-sdk-lmi-e2e" +DEFAULT_STACK = "java-lmi-e2e" +FIXTURES = {"default1": (1, "default"), "default2": (2, "default"), + "default8": (8, "default"), "fixed2": (2, "fixed"), "nested2": (2, "fixed")} + + +def template(manifest): + resources, outputs = {}, {} + for key, (concurrency, executor) in FIXTURES.items(): + name = manifest["stack"] + "-" + key + log_id, fn_id = key + "Logs", key + "Function" + resources[log_id] = {"Type": "AWS::Logs::LogGroup", "Properties": { + "LogGroupName": "/aws/lambda/" + name, "RetentionInDays": 1}} + resources[fn_id] = {"Type": "AWS::Lambda::Function", "Properties": { + "FunctionName": name, "Runtime": "java25", "Architectures": ["arm64"], + "Role": manifest["role"], "Handler": "software.amazon.lambda.durable.lmi.LifecycleHandler", + "Code": {"S3Bucket": manifest["bucket"], "S3Key": manifest["codeKey"]}, + "Timeout": manifest["invocationTimeout"], "MemorySize": 2048, + "FunctionScalingConfig": {"MinExecutionEnvironments": 1, "MaxExecutionEnvironments": 1}, + "DurableConfig": {"ExecutionTimeout": 240, "RetentionPeriodInDays": 1}, + "CapacityProviderConfig": {"LambdaManagedInstancesCapacityProviderConfig": { + "CapacityProviderArn": manifest["provider"], + "PerExecutionEnvironmentMaxConcurrency": concurrency, + "ExecutionEnvironmentMemoryGiBPerVCpu": 2}}, + "Environment": {"Variables": {"LMI_EXECUTOR": executor, "LMI_COMMIT": manifest["commit"], + "LMI_TEST_RUN_ID": manifest["runId"]}}, + "LoggingConfig": {"LogFormat": "JSON", "ApplicationLogLevel": "INFO", + "SystemLogLevel": "INFO", "LogGroup": {"Ref": log_id}}}} + # CloudFormation automatically publishes $LATEST.PUBLISHED for an LMI function. + # An additional numbered version would provision a second independent set of environments. + outputs[key] = {"Value": {"Fn::Join": ["", [{"Fn::GetAtt": [fn_id, "Arn"]}, ":$LATEST.PUBLISHED"]]}} + return {"AWSTemplateFormatVersion": "2010-09-09", "Resources": resources, "Outputs": outputs} + + +def deploy(run_id, invocation_timeout, stack_name=DEFAULT_STACK): + if not re.fullmatch(r"[a-z0-9-]{1,24}", run_id): + raise PreconditionError("run-id must be 1-24 lowercase letters, digits, or hyphens") + if not re.fullmatch(r"[A-Za-z][A-Za-z0-9-]{0,49}", stack_name): + raise PreconditionError("stack-name must be 1-50 letters, digits, or hyphens, starting with a letter") + provider = os.environ["CAPACITY_PROVIDER_ARN"] + region = provider.split(":")[3] + if region != os.environ.get("AWS_REGION"): + raise PreconditionError("AWS_REGION must match the capacity provider region") + account = aws("sts", "get-caller-identity")["Account"] + if account != provider.split(":")[4]: + raise PreconditionError("Capacity provider must belong to the authenticated test account") + name = provider.rsplit(":", 1)[-1].rsplit("/", 1)[-1] + capacity = aws("lambda", "get-capacity-provider", {"CapacityProviderName": name}) + save(ARTIFACTS / "capacity-provider.json", capacity) + scaling = capacity["CapacityProvider"].get("CapacityProviderScalingConfig", {}) + if not 2 <= scaling.get("MaxVCpuCount", 0) <= 128: + raise PreconditionError("Dedicated provider must have an explicit maximum of 2-128 vCPUs") + jar = ROOT / "target/lmi-fixtures.jar" + digest = hashlib.sha256(jar.read_bytes()).digest() + bucket_scope = hashlib.sha256(f"{region}:{stack_name}".encode()).hexdigest()[:12] + manifest = {"runId": run_id, "stack": stack_name, "persistent": True, "account": account, + "bucket": f"java-lmi-e2e-{account}-{bucket_scope}", "region": region, + "role": os.environ["TEST_LAMBDA_EXECUTION_ROLE_ARN"], "provider": provider, + "commit": subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), + "codeSha256": base64.b64encode(digest).decode(), "codeKey": f"code/{digest.hex()}.jar", + "invocationTimeout": invocation_timeout, "created": int(time.time()), "functions": {}} + save(MANIFEST, manifest) + tags = [{"Key": "Suite", "Value": OWNER}, {"Key": "Stack", "Value": stack_name}, + {"Key": "Persistent", "Value": "true"}] + ensure_bucket(manifest, tags) + aws("s3api", "put-object", {"Bucket": manifest["bucket"], "Key": manifest["codeKey"]}, extra=["--body", str(jar)]) + deploy_fixtures(manifest, tags) + manifest["logStartMillis"] = int(time.time() * 1000) + save(MANIFEST, manifest) + + +def ensure_bucket(manifest, tags): + """Reuse the suite-owned bucket; only per-run control objects have automatic expiry.""" + bucket = manifest["bucket"] + exists = True + try: + aws("s3api", "head-bucket", {"Bucket": bucket, "ExpectedBucketOwner": manifest["account"]}) + except RuntimeError as error: + if not any(text in str(error) for text in ("404", "Not Found", "NoSuchBucket")): + raise + exists = False + if exists: + actual = aws("s3api", "get-bucket-tagging", {"Bucket": bucket})["TagSet"] + owner = {tag["Key"]: tag["Value"] for tag in actual} + if owner.get("Suite") != OWNER or owner.get("Stack") != manifest["stack"]: + raise PreconditionError("The persistent bucket is not owned by this test stack") + else: + request = {"Bucket": bucket} + if manifest["region"] != "us-east-1": + request["CreateBucketConfiguration"] = {"LocationConstraint": manifest["region"]} + aws("s3api", "create-bucket", request) + aws("s3api", "put-bucket-tagging", {"Bucket": bucket, "Tagging": {"TagSet": tags}}) + aws("s3api", "put-public-access-block", {"Bucket": bucket, "PublicAccessBlockConfiguration": { + "BlockPublicAcls": True, "IgnorePublicAcls": True, "BlockPublicPolicy": True, "RestrictPublicBuckets": True}}) + aws("s3api", "put-bucket-lifecycle-configuration", {"Bucket": bucket, + "LifecycleConfiguration": {"Rules": [{"ID": "expire-controls", "Status": "Enabled", + "Filter": {"Prefix": "control/"}, "Expiration": {"Days": 1}, + "AbortIncompleteMultipartUpload": {"DaysAfterInitiation": 1}}]}}) + + +def remaining_budget(deadline, maximum): + remaining = int(deadline - time.monotonic()) + if remaining <= 0: + raise PreconditionError("Five-function provisioning budget exhausted") + return min(remaining, maximum) + + +def existing_stack(name): + try: + return aws("cloudformation", "describe-stacks", {"StackName": name})["Stacks"][0] + except RuntimeError as error: + if "does not exist" in str(error): + return None + raise + + +def wait_for_idle_functions(stack, seconds=270): + """Do not change code while previous executions of these persistent fixtures are running.""" + deadline = time.monotonic() + seconds + while True: + running = [] + for output in stack.get("Outputs", []): + if output["OutputKey"] not in FIXTURES: + continue + function = output["OutputValue"].rsplit(":", 1)[0] + marker = None + while True: + request = {"FunctionName": function, "Statuses": ["RUNNING"]} + if marker: + request["Marker"] = marker + response = aws("lambda", "list-durable-executions-by-function", request, extra=["--no-paginate"]) + running.extend(response.get("DurableExecutions", [])) + marker = response.get("NextMarker") + if not marker: + break + if time.monotonic() >= deadline: + raise PreconditionError("Could not finish checking previous executions; persistent code was not updated") + save(ARTIFACTS / "pre-deploy-executions.json", running) + if not running: + return + if time.monotonic() >= deadline: + raise PreconditionError("Previous durable test executions are still running; persistent code was not updated") + time.sleep(2) + + +def deploy_fixtures(manifest, tags, seconds=1800): + """Create once, then update the same stack and functions without test teardown.""" + deadline = time.monotonic() + seconds + remaining_budget(deadline, seconds) + save(MANIFEST, manifest) + spec = template(manifest) + save(ARTIFACTS / "template.json", spec) + stack = existing_stack(manifest["stack"]) + request = {"StackName": manifest["stack"], "TemplateBody": json.dumps(spec), "Tags": tags} + if stack is None: + request["TimeoutInMinutes"] = 25 + aws("cloudformation", "create-stack", request) + expected = "CREATE_COMPLETE" + else: + owner = {tag["Key"]: tag["Value"] for tag in stack.get("Tags", [])} + if owner.get("Suite") != OWNER: + raise PreconditionError("The persistent stack is not owned by this suite") + if stack["StackStatus"] not in {"CREATE_COMPLETE", "UPDATE_COMPLETE", "UPDATE_ROLLBACK_COMPLETE"}: + raise PreconditionError(f"Persistent stack requires recovery from {stack['StackStatus']}; it was retained") + wait_for_idle_functions(stack, seconds=remaining_budget(deadline, 270)) + try: + aws("cloudformation", "update-stack", request) + expected = "UPDATE_COMPLETE" + except RuntimeError as error: + if "No updates are to be performed" not in str(error): + raise + expected = None + if expected: + wait_stack(manifest["stack"], expected, remaining_budget(deadline, seconds)) + stack = existing_stack(manifest["stack"]) + outputs = {output["OutputKey"]: output["OutputValue"] for output in stack["Outputs"]} + require(set(outputs) == set(FIXTURES), "The stack must expose all five LMI fixtures") + for fixture in FIXTURES: + arn = outputs[fixture] + record_fixture(manifest, fixture, arn) + verify_function_scaling(arn, fixture, seconds=remaining_budget(deadline, 300)) + + +def record_fixture(manifest, key, arn): + config = aws("lambda", "get-function-configuration", {"FunctionName": arn}) + save(ARTIFACTS / "configuration" / (key + ".json"), config) + actual = config.get("CapacityProviderConfig", {}).get("LambdaManagedInstancesCapacityProviderConfig", {}) + require(actual.get("CapacityProviderArn") == manifest["provider"], "Deployment is not associated with the requested LMI provider") + require(actual.get("PerExecutionEnvironmentMaxConcurrency") == FIXTURES[key][0], "Concurrency readback mismatch") + require(config["Runtime"] == "java25" and config["Architectures"] == ["arm64"], "Unsupported runtime/architecture") + require(config.get("DurableConfig", {}).get("ExecutionTimeout") == 240, "Function is not durable") + require(config["Version"] == "$LATEST.PUBLISHED" and config["CodeSha256"] == manifest["codeSha256"], "Artifact/version mismatch") + require(config["MemorySize"] == 2048, "Fixture must use the minimum 2 GiB / 1 vCPU allocation") + manifest["functions"][key] = {"arn": arn, "logGroup": config["LoggingConfig"]["LogGroup"], "concurrency": FIXTURES[key][0]} + save(MANIFEST, manifest) + + +def verify_function_scaling(arn, fixture, seconds=300): + """Read back CloudFormation's applied limits and wait for the published version to be active.""" + desired = {"MinExecutionEnvironments": 1, "MaxExecutionEnvironments": 1} + function_name, qualifier = arn.rsplit(":", 1) + request = {"FunctionName": function_name, "Qualifier": qualifier} + deadline = time.monotonic() + seconds + while time.monotonic() < deadline: + scaling = aws("lambda", "get-function-scaling-config", request) + config = aws("lambda", "get-function-configuration", {"FunctionName": arn}) + save(ARTIFACTS / "configuration" / (fixture + "-scaling.json"), scaling) + save(ARTIFACTS / "configuration" / (fixture + ".json"), config) + if config.get("State") == "Failed": + raise PreconditionError("LMI version provisioning failed: " + config.get("StateReason", "unknown")) + if scaling.get("AppliedFunctionScalingConfig") == desired and config.get("State") == "Active": + return + time.sleep(2) + raise PreconditionError("Function scaling/provisioning did not reach the bounded configuration") + + +def wait_stack(name, expected, seconds): + deadline = time.monotonic() + seconds + while time.monotonic() < deadline: + response = aws("cloudformation", "describe-stacks", {"StackName": name}) + status = response["Stacks"][0]["StackStatus"] + if status == expected: + return + if "FAILED" in status or "ROLLBACK" in status: + save(ARTIFACTS / "stack-events.json", aws("cloudformation", "describe-stack-events", {"StackName": name})) + raise PreconditionError(f"Deployment {status}: see stack-events.json; no fallback to ordinary Lambda") + time.sleep(5) + raise PreconditionError("Provisioning budget exhausted") + + +def wait_returns(cloud, fixture, items, seconds=30): + markers = {i["marker"] for i in items} + cloud.poll(fixture, lambda events: markers <= {e["marker"] for e in selected(events, "WRAPPER_RETURN")}, seconds, items=items) + + +def healthy_peers(cloud, fixture, target, count, prefix): + """Bounded placement only; never retry an established lifecycle assertion.""" + gate_name = prefix + "-gate" + gate = cloud.gate(gate_name) + admitted, attempts = [], [] + deadline = time.monotonic() + 25 + for attempt in range(4): + if len(admitted) >= count or time.monotonic() >= deadline: + break + batch = [cloud.launch(fixture, "hold", f"{prefix}-{attempt}-{i}", target=target, gate=gate) + for i in range(count - len(admitted))] + attempts.extend(batch) + markers = {i["marker"] for i in batch} + cloud.poll(fixture, lambda events: markers <= {e["marker"] for e in events + if e["kind"] in {"HEARTBEAT", "PLACEMENT_MISS"}}, seconds=8, category=PreconditionError, items=batch) + admitted += [i for i in batch if selected(cloud.events_for(i), "HEARTBEAT")] + if len(admitted) != count: + cloud.gate(gate_name, release=True) + raise PreconditionError(f"Could not place {count} healthy peers in original environment {target}") + return admitted, gate_name + + +def replay_case(cloud, fixture, scenario): + prefix = uuid.uuid4().hex[:12] + # A healthy holder supplies the target JVM before the suspension is triggered. + anchor = None + if fixture != "default1": + gate_name = prefix + "-anchor" + gate = cloud.gate(gate_name) + anchor = cloud.launch(fixture, "hold", prefix + "-healthy", gate=gate) + evidence = cloud.poll(fixture, lambda events: selected(events, "HEARTBEAT", anchor["marker"]), category=PreconditionError, items=[anchor]) + target = evidence[0]["environment"] + else: + target = None + victim = None + for attempt in range(4): + item = cloud.launch(fixture, scenario, f"{prefix}-victim-{attempt}", target=target) + cloud.poll(fixture, lambda events: selected(events, "WRAPPER_RETURN", item["marker"]), seconds=30, items=[item]) + if not selected(cloud.events_for(item), "PLACEMENT_MISS"): + victim = item + break + if victim is None: + raise PreconditionError("No replay victim admitted to the anchor JVM") + cloud.finish(victim) + cloud.poll(fixture, lambda events: any(e.get("status") == "SUCCEEDED" for e in selected(events, "WRAPPER_RETURN", victim["marker"]))) + if anchor: + cloud.gate(gate_name, release=True) + cloud.finish(anchor) + wait_returns(cloud, fixture, [anchor]) + # Root/task interval overlap is required, including the first victim invocation. + assert_overlap(list(cloud.events.values()), {anchor["marker"], victim["marker"]}, 2, target) + history = cloud.history(victim) + events = cloud.events_for(victim) + assert_replay(events, history, victim["marker"]) + if scenario == "suspend": + if anchor: + cleanup = selected(events, "CLEANUP_ENTER")[0] + exits = selected(events, "CLEANUP_EXIT") + require(any(e["environment"] == target and cleanup["nanos"] <= e["nanos"] <= exits[0]["nanos"] + for e in selected(cloud.events_for(anchor), "HEARTBEAT")), + "Healthy invocation did not progress during root cleanup") + assert_lifecycle(events) + + +def overlap_case(cloud, fixture, target=None): + count = FIXTURES[fixture][0] + prefix = uuid.uuid4().hex[:12] + gate_name = prefix + "-anchor" + anchor = cloud.launch(fixture, "hold", prefix + "-anchor", target=target, gate=cloud.gate(gate_name)) + heartbeat = cloud.poll(fixture, lambda events: selected(events, "HEARTBEAT", anchor["marker"]), category=PreconditionError, items=[anchor])[0] + peers, peer_gate = healthy_peers(cloud, fixture, heartbeat["environment"], count - 1, prefix + "-peer") + items = [anchor] + peers + assert_overlap(list(cloud.events.values()), {i["marker"] for i in items}, count, heartbeat["environment"]) + cloud.gate(gate_name, release=True) + cloud.gate(peer_gate, release=True) + for item in items: + cloud.finish(item) + cloud.history(item) + wait_returns(cloud, fixture, items) + for item in items: + assert_lifecycle(cloud.events_for(item)) + + +def fixed_case(cloud, fixture, scenario): + for attempt in range(3): + prefix = uuid.uuid4().hex[:12] + items = [cloud.launch(fixture, scenario, prefix + f"-{i}", cohort=prefix, peers=2) for i in range(2)] + wait_returns(cloud, fixture, items, seconds=35) + events = [e for i in items for e in cloud.events_for(i)] + if len(selected(events, "BARRIER_PASSED")) == 2: + # A failure after admission is final, even if the escape grows the pool. + for item in items: + cloud.finish(item) + cloud.history(item) + assert_fixed(events, {i["marker"] for i in items}) + assert_lifecycle(events) + return + if selected(events, "BARRIER_PASSED") or selected(events, "ESCAPE"): + raise AssertionError("Partial fixed-pool admission or escape; refusing to rerun the assertion") + raise PreconditionError("Two fixed-executor roots could not meet in one JVM") + + +def invocation_deadline(trace): + entry = selected(trace, "WRAPPER_ENTER")[0] + return entry["nanos"] + entry["remainingMillis"] * 1_000_000 + + +def timeout_case(cloud, fixture, stubborn=False): + prefix = uuid.uuid4().hex[:12] + timeout = cloud.manifest["invocationTimeout"] + scenario = "stubborn" if stubborn else "timeout" + victim = cloud.launch(fixture, scenario, prefix + "-victim", hold_ms=(timeout + 20) * 1000) + entry = cloud.poll(fixture, lambda events: selected(events, "TASK_ENTER", victim["marker"]), category=PreconditionError, items=[victim])[0] + target = entry["environment"] + runtime_entry = selected(cloud.events_for(victim), "WRAPPER_ENTER")[0] + deadline_wall = (runtime_entry["epochMillis"] + runtime_entry["remainingMillis"]) / 1000 + # Stagger admission: healthy invocations must outlive the victim's real deadline. + cloud.poll(fixture, lambda events: time.time() >= deadline_wall - timeout / 2, seconds=timeout) + peers, gate = healthy_peers(cloud, fixture, target, FIXTURES[fixture][0] - 1, prefix + "-peer") + assert_overlap(list(cloud.events.values()), {victim["marker"], *(p["marker"] for p in peers)}, FIXTURES[fixture][0], target) + deadline = invocation_deadline(cloud.events_for(victim)) + # Launch probes while all other established slots remain occupied. Retry placement only. + cloud.poll(fixture, lambda events: time.time() >= deadline_wall - 5, seconds=timeout) + probe_gate_name = prefix + "-probe-gate" + probe_gate = cloud.gate(probe_gate_name) + probes = [cloud.launch(fixture, "probe", prefix + f"-probe-{i}", target=target, gate=probe_gate) for i in range(4)] + cloud.poll(fixture, lambda events: time.time() >= deadline_wall + 8, seconds=20) + cloud.gate(gate, release=True) + cloud.gate(probe_gate_name, release=True) + for peer in peers: + cloud.finish(peer) + # Residual Java code is bounded, but arbitrary code cannot be forcibly killed. + wait_returns(cloud, fixture, [victim], seconds=timeout + 30) + cloud.finish(victim, expected=None) + history = cloud.history(victim) + trace = cloud.events_for(victim) + returned = selected(trace, "WRAPPER_RETURN")[0] + recovered = [e for p in probes for e in selected(cloud.events_for(p), "TASK_ENTER") + if e["environment"] == target and e["nanos"] <= deadline + 8_000_000_000] + report = {"requestId": entry["requestId"], "environment": target, "deadlineNanos": deadline, + "wrapperReturnNanos": returned["nanos"], "taskExits": selected(trace, "TASK_EXIT"), + "interrupts": selected(trace, "INTERRUPTED") + selected(trace, "IGNORED_INTERRUPT"), + "recoveryAdmissions": recovered, "durableHistoryEvents": len(history)} + report["serverTimeoutLogs"] = [e for e in cloud.raw_logs.values() if entry["requestId"] in e["message"] + and ("timeout" in e["message"].lower() or "timed out" in e["message"].lower())] + save(ARTIFACTS / "timeouts" / (prefix + ".json"), report) + errors = [] + def check(condition, message): + if not condition: + errors.append(message) + check(returned["nanos"] <= deadline + 5_000_000_000, "SDK cleanup exceeded the actual invocation deadline + 5s") + check(bool(report["serverTimeoutLogs"]) or (returned["status"] == "THREW" and bool(report["interrupts"])), + "Neither server-reported invocation timeout nor explicit SDK deadline cancellation observed") + check(bool(recovered), "Affected worker slot was not recovered in the original JVM while healthy slots stayed occupied") + for peer in peers: + heartbeats = selected(cloud.events_for(peer), "HEARTBEAT") + check(any(e["nanos"] >= deadline for e in heartbeats), "Healthy invocation stopped progressing at victim deadline") + assert_lifecycle(cloud.events_for(peer)) + if stubborn: + check(bool(selected(trace, "IGNORED_INTERRUPT")), "Cancellation was not attempted for the non-cooperative child") + check(bool(selected(trace, "LATE_REJECTED")) and not selected(trace, "LATE_ACCEPTED"), "Residual child could issue later SDK work") + check(not [e for e in selected(trace, "BODY") if e["name"] == "late-work"], "Late step body executed") + else: + check(bool(selected(trace, "INTERRUPTED")), "Interruptible invocation task was not cancelled") + exits = selected(trace, "TASK_EXIT") + check(bool(exits) and exits[0]["nanos"] <= deadline + 5_000_000_000, "Task exit exceeded cleanup budget") + check(not selected(trace, "ESCAPE"), "Timed-out task survived until the test-only escape") + if errors: + raise AssertionError("; ".join(errors)) + assert_lifecycle(trace, allow_residual=stubborn) + # All lanes, not just one replacement invocation, must be available again. + overlap_case(cloud, fixture, target=target) + + +def inflight_case(cloud, fixture, scenario): + prefix = uuid.uuid4().hex[:12] + gate_name = prefix + "-gate" + anchor = cloud.launch(fixture, "hold", prefix + "-healthy", gate=cloud.gate(gate_name)) + target = cloud.poll(fixture, lambda events: selected(events, "HEARTBEAT", anchor["marker"]), category=PreconditionError, items=[anchor])[0]["environment"] + victim = cloud.launch(fixture, scenario, prefix + "-victim", target=target, hold_ms=1500) + wait_returns(cloud, fixture, [victim]) + cloud.gate(gate_name, release=True) + cloud.finish(anchor) + if selected(cloud.events_for(victim), "PLACEMENT_MISS"): + raise PreconditionError("In-flight cleanup victim placed in another environment") + cloud.finish(victim, expected="FAILED" if scenario == "failure-inflight" else "SUCCEEDED") + cloud.history(victim) + assert_overlap(list(cloud.events.values()), {anchor["marker"], victim["marker"]}, 2, target) + assert_lifecycle(cloud.events_for(victim)) + + +def warm_case(cloud, fixture): + target, snapshots = None, [] + for batch in range(3): + for scenario in ["success", "failure", "replay"]: + marker = uuid.uuid4().hex[:12] + item = cloud.launch(fixture, scenario, marker, target=target) + wait_returns(cloud, fixture, [item]) + if selected(cloud.events_for(item), "PLACEMENT_MISS"): + raise PreconditionError("Warm-environment fixture replaced; cannot claim no accumulation") + cloud.finish(item, expected="FAILED" if scenario == "failure" else "SUCCEEDED") + trace = cloud.events_for(item) + if target is None: + target = trace[0]["environment"] + snapshots.extend(e for e in selected(trace, "SNAPSHOT") if e["environment"] == target) + if scenario == "replay": + assert_replay(trace, cloud.history(item), marker) + assert_lifecycle(trace) + # Default cached pools retain idle threads: measure live invocation-owned work/queue instead. + require(snapshots, "Missing warm task snapshots") + require(all(e["liveTasks"] == 0 and e["liveRoots"] == 0 and e["queued"] == 0 for e in snapshots), + "Invocation-owned tasks, roots or queued work accumulated across warm batches") + require(snapshots[-1]["threads"] <= snapshots[0]["threads"] + 16, + "Thread growth exceeded bounded cache tolerance (+16)") + + +def cases_for_fixture(cloud, fixture): + if fixture == "default1": + return [("baseline-concurrency1", lambda: replay_case(cloud, fixture, "baseline"))] + if fixture in {"fixed2", "nested2"}: + scenario = "fixed" if fixture == "fixed2" else "nested" + return [("fixed-two-roots" if fixture == "fixed2" else "fixed-nested-map-parallel", + lambda: fixed_case(cloud, fixture, scenario))] + cases = [(fixture + "-isolation", lambda: overlap_case(cloud, fixture)), + (fixture + "-suspend-cleanup-replay", lambda: replay_case(cloud, fixture, "suspend")), + (fixture + "-timeout-recovery", lambda: timeout_case(cloud, fixture))] + if fixture == "default2": + cases += [("non-cooperative-child", lambda: timeout_case(cloud, fixture, True)), + ("return-with-inflight-step", lambda: inflight_case(cloud, fixture, "return-inflight")), + ("failure-with-inflight-step", lambda: inflight_case(cloud, fixture, "failure-inflight")), + ("warm-repeated-batches", lambda: warm_case(cloud, fixture))] + return cases + + +def run_tests(): + manifest = json.loads(MANIFEST.read_text()) + if set(manifest["functions"]) != set(FIXTURES): + raise PreconditionError("All five functions must be deployed before running the cloud suite") + cloud = Cloud(manifest, ARTIFACTS) + suite = ET.Element("testsuite", name="LMI cloud lifecycle") + cases = [case for fixture in FIXTURES for case in cases_for_fixture(cloud, fixture)] + try: + for name, case in cases: + started = time.monotonic() + node = ET.SubElement(suite, "testcase", classname="LmiCloudLifecycle", name=name) + try: + case() + print("PASS", name, flush=True) + except Exception as error: + kind = "failure" if isinstance(error, AssertionError) else "error" + ET.SubElement(node, kind, type=type(error).__name__, message=str(error)).text = traceback.format_exc() + print(kind.upper(), name, str(error), flush=True) + finally: + try: + cloud.release_all() + for fixture in manifest["functions"]: + cloud.refresh(fixture) + except Exception as error: + ET.SubElement(node, "error", type="CollectionError", message=str(error)) + node.set("time", str(round(time.monotonic() - started, 3))) + write_junit(suite) + finally: + cloud.close() + write_junit(suite) + require(not suite.findall(".//failure") and not suite.findall(".//error"), "LMI cloud cases failed; see JUnit and artifacts") + + +def write_junit(suite): + suite.set("tests", str(len(suite.findall("testcase")))) + suite.set("failures", str(len(suite.findall(".//failure")))) + suite.set("errors", str(len(suite.findall(".//error")))) + ARTIFACTS.mkdir(parents=True, exist_ok=True) + ET.ElementTree(suite).write(ARTIFACTS / "junit.xml", encoding="unicode", xml_declaration=True) + + +def collect(): + manifest = json.loads(MANIFEST.read_text()) + cloud = Cloud(manifest, ARTIFACTS) + errors = [] + for fixture in manifest["functions"]: + try: + cloud.refresh(fixture) + except Exception as error: + errors.append(str(error)) + executions = {e["executionArn"]: e["marker"] for e in cloud.events.values()} + for arn, marker in executions.items(): + try: + cloud.history({"arn": arn, "marker": marker}) + save(ARTIFACTS / "executions" / (marker + ".json"), aws("lambda", "get-durable-execution", {"DurableExecutionArn": arn})) + except Exception as error: + errors.append(str(error)) + cloud.pool.shutdown() + save(ARTIFACTS / "collection-errors.json", errors) + require(not errors, "Evidence collection failed: " + "; ".join(errors)) + + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("command", choices=["deploy", "test", "collect"]) + parser.add_argument("--run-id") + parser.add_argument("--stack-name", default=DEFAULT_STACK) + parser.add_argument("--cloud-enabled", action="store_true") + parser.add_argument("--invocation-timeout", type=int, default=60, choices=range(45, 91)) + args = parser.parse_args() + try: + if args.command == "deploy": + deploy(args.run_id, args.invocation_timeout, args.stack_name) + elif args.command == "test": + if not args.cloud_enabled: + parser.error("Real cloud tests require --cloud-enabled") + run_tests() + else: + collect() + except Exception as error: + save(ARTIFACTS / (args.command + "-error.json"), {"category": args.command, "error": str(error)}) + raise + + +if __name__ == "__main__": + main() diff --git a/lmi-tests/cloud_support.py b/lmi-tests/cloud_support.py new file mode 100644 index 000000000..30a684549 --- /dev/null +++ b/lmi-tests/cloud_support.py @@ -0,0 +1,295 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Cloud API adapter and evidence assertions; no third-party Python dependencies.""" +import concurrent.futures +import json +import os +from pathlib import Path +import re +import subprocess +import tempfile +import time + + +class PreconditionError(RuntimeError): + """The scenario did not establish the required infrastructure/placement.""" + + +class CollectionError(RuntimeError): + """Required cloud evidence could not be retrieved.""" + + +def require(condition, message): + if not condition: + raise AssertionError(message) + + +def scrub(value): + if isinstance(value, dict): + return {k: ("" if k in {"controlUrl", "CheckpointToken"} else scrub(v)) + for k, v in value.items()} + if isinstance(value, list): + return [scrub(v) for v in value] + if isinstance(value, str): + return re.sub(r'https://[^\s"<>]+', '', value) + return value + + +def save(path, value): + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(scrub(value), indent=2, default=str) + "\n") + + +def aws(service, operation, data=None, extra=(), timeout=30, raw=False): + command = ["aws", service, operation, "--no-cli-pager", "--output", "json", + "--cli-connect-timeout", "5", "--cli-read-timeout", str(timeout)] + if data is not None: + command += ["--cli-input-json", json.dumps(data)] + command += list(extra) + result = subprocess.run(command, capture_output=True, text=True, timeout=timeout + 10, + env={**os.environ, "AWS_MAX_ATTEMPTS": "2", "AWS_PAGER": ""}) + if result.returncode: + raise RuntimeError(f"{service} {operation} (exit {result.returncode}): {scrub(result.stderr[-4000:])}") + return result.stdout.strip() if raw else json.loads(result.stdout or "{}") + + +def diagnostic(message): + """Accept both raw stdout and the LMI structured JSON logging envelope.""" + try: + envelope = json.loads(message) + if isinstance(envelope, dict): + message = envelope.get("message", "") + except (ValueError, TypeError): + pass + if not isinstance(message, str) or "LMI_TEST " not in message: + return None + try: + return json.loads(message.split("LMI_TEST ", 1)[1]) + except ValueError: + return None + + +def selected(events, kind=None, marker=None): + return [e for e in events if (kind is None or e["kind"] == kind) + and (marker is None or e["marker"] == marker)] + + +def assert_overlap(events, markers, count, environment=None): + """Use actual task intervals on one JVM; driver parallelism is not evidence.""" + points = {} + for event in events: + if event["marker"] not in markers or event["kind"] not in {"TASK_ENTER", "TASK_EXIT"}: + continue + if environment is not None and event["environment"] != environment: + continue + points.setdefault(event["environment"], []).append(event) + for env, entries in points.items(): + active = {} + for event in sorted(entries, key=lambda e: e["sequence"]): + key = event["requestId"] + active[key] = active.get(key, 0) + (1 if event["kind"] == "TASK_ENTER" else -1) + if sum(n > 0 for n in active.values()) >= count: + return env + raise PreconditionError(f"No evidence of {count} overlapping request IDs in one JVM") + + +def assert_lifecycle(events, allow_residual=False): + returns = selected(events, "WRAPPER_RETURN") + require(returns, "No SDK wrapper return observed") + for returned in returns: + local = [e for e in events if e["requestId"] == returned["requestId"] + and e["environment"] == returned["environment"]] + if returned["status"] in {"SUCCEEDED", "FAILED", "PENDING"}: + require(returned["rootExited"], f"{returned['status']} returned before root exit") + require(returned["tasks"] == 0, f"{returned['status']} returned with live invocation tasks") + roots = selected(local, "ROOT_EXIT") + require(roots and roots[-1]["sequence"] < returned["sequence"], "Missing causal root exit") + if not allow_residual: + require(not selected(local, "ESCAPE"), "A test-only escape was needed to finish SDK work") + late = [e for e in local if e["sequence"] > returned["sequence"] + and e["kind"] in {"CHECKPOINT_CALL", "POLL_CALL", "CHECKPOINT_EXIT", "POLL_EXIT"}] + require(not late, "SDK checkpoint/poll activity continued after wrapper return") + + +def assert_replay(events, history, marker): + calls = selected(events, "WRAPPER_ENTER", marker) + require(len({e["requestId"] for e in calls}) >= 2, "No real invocation resume observed") + require(any(e.get("status") == "PENDING" for e in selected(events, "WRAPPER_RETURN", marker)), + "No real suspension observed") + for name, event_type in [("success", "StepSucceeded"), ("failure", "StepFailed")]: + bodies = [e for e in selected(events, "BODY", marker) if e["name"] == name] + require(len(bodies) == 1, f"Checkpointed {name} body ran {len(bodies)} times") + require(bodies[0]["value"] == marker, "Cross-execution result contamination") + entries = [e for e in history if e.get("Name") == name] + terminal = [e for e in entries if e.get("EventType") == event_type] + require(len(terminal) == 1, f"Missing/duplicate real {event_type} history") + require(len({e["Id"] for e in entries}) == 1, f"{name} operation identity changed") + failures = selected(events, "STORED_FAILURE", marker) + require(len(failures) >= 2 and all(e["message"] == "expected:" + marker for e in failures), + "Stored failure meaning was not reproduced on replay") + require(any(e.get("EventType") == "WaitSucceeded" for e in history), "No service wait completion") + + +def assert_fixed(events, markers): + entered = selected(events, "BARRIER_ENTER") + passed = selected(events, "BARRIER_PASSED") + require(len({e["marker"] for e in passed}) == len(markers), "Root barrier was not passed by every participant") + envs = {e["environment"] for e in passed} + if len(envs) != 1: + raise PreconditionError("Fixed-executor roots were placed in different JVMs") + require(len({e["requestId"] for e in entered}) == len(markers), "Missing distinct runtime invocations") + require(max(e["sequence"] for e in entered) < min(e["sequence"] for e in passed), + "Root invocations did not overlap at the barrier") + require(not selected(events, "ESCAPE"), "Shared fixed executor starved its own queued work") + for marker in markers: + progress = selected(events, "PROGRESS", marker) + start = selected(events, "BARRIER_PASSED", marker) + require(progress and (progress[-1]["nanos"] - start[0]["nanos"]) < 8_000_000_000, + "Shared executor did not progress within budget") + + +class Cloud: + def __init__(self, manifest, artifacts): + self.manifest, self.artifacts = manifest, Path(artifacts) + self.pool = concurrent.futures.ThreadPoolExecutor(max_workers=32) + self.events, self.raw_logs, self.invocations = {}, {}, [] + self.start_ms = manifest.get("logStartMillis", int(time.time() * 1000)) + self.gates = set() + + def gate(self, name, release=False): + self.gates.add(name) + with tempfile.NamedTemporaryFile(mode="w") as body: + body.write("release" if release else "hold") + body.flush() + aws("s3api", "put-object", {"Bucket": self.manifest["bucket"], "Key": "control/" + self.manifest["runId"] + "/" + name}, + extra=["--body", body.name]) + if release: + return None + return aws("s3", "presign", extra=[f"s3://{self.manifest['bucket']}/control/{self.manifest['runId']}/{name}", + "--expires-in", "3600"], raw=True) + + def launch(self, fixture, scenario, marker, cohort=None, target=None, peers=1, hold_ms=100000, gate=None): + item = {"marker": marker, "fixture": fixture, "started": time.time()} + payload = {"runId": self.manifest["runId"], "cohort": cohort or marker, + "scenario": scenario, "marker": marker, "controlUrl": gate, + "targetEnvironment": target, "peers": peers, "holdMillis": hold_ms} + item["future"] = self.pool.submit(self._invoke, fixture, payload) + self.invocations.append(item) + return item + + def _invoke(self, fixture, payload): + started = time.time() + artifact = self.artifacts / "invocations" / (payload["marker"] + ".json") + details = {"fixture": fixture, "functionArn": self.manifest["functions"][fixture]["arn"], + "scenario": payload["scenario"], "marker": payload["marker"], "started": started} + save(artifact, {**details, "state": "STARTED"}) + try: + result = self._invoke_request(fixture, payload) + save(artifact, {**details, **result, "state": "RETURNED", "elapsedSeconds": time.time() - started}) + return result + except Exception as error: + save(artifact, {**details, "state": "REQUEST_FAILED", "elapsedSeconds": time.time() - started, + "errorType": type(error).__name__, "error": str(error)}) + raise + + def _invoke_request(self, fixture, payload): + with tempfile.TemporaryDirectory() as directory: + path = Path(directory) / "response.json" + source = Path(directory) / "payload.json" + source.write_text(json.dumps(payload), encoding="utf-8") + # lambda invoke is a custom AWS CLI command: required flags cannot come from --cli-input-json. + # fileb:// sends the original JSON bytes regardless of the user's CLI binary-format setting. + headers = aws("lambda", "invoke", extra=[ + "--function-name", self.manifest["functions"][fixture]["arn"], + "--invocation-type", "Event" if payload["scenario"] in {"timeout", "stubborn"} else "RequestResponse", + "--payload", "fileb://" + str(source), str(path)], timeout=150) + try: + body = json.loads(path.read_text()) + except ValueError: + body = path.read_text() + return {"headers": headers, "body": body} + + def refresh(self, fixture): + group = self.manifest["functions"][fixture]["logGroup"] + data = aws("logs", "filter-log-events", {"logGroupName": group, "startTime": self.start_ms}) + for event in data.get("events", []): + self.raw_logs[event["eventId"]] = event + parsed = diagnostic(event["message"]) + if parsed and parsed.get("runId") == self.manifest["runId"]: + parsed["fixture"] = fixture + self.events[(parsed["environment"], parsed["sequence"])] = parsed + save(self.artifacts / "diagnostics.json", list(self.events.values())) + save(self.artifacts / "cloudwatch.json", list(self.raw_logs.values())) + events = [e for e in self.events.values() if e["fixture"] == fixture] + if any(e.get("deploymentRunId") != self.manifest["runId"] or e.get("commit") != self.manifest["commit"] + for e in events): + raise CollectionError("Invocation reached an outdated deployment; inspect the recorded commit and deploymentRunId") + return events + + def poll(self, fixture, predicate, seconds=25, category=AssertionError, items=()): + deadline = time.monotonic() + seconds + while True: + for item in items: + future = item["future"] + if future.done() and future.exception() is not None: + error = future.exception() + raise CollectionError(f"Invocation request failed for {item['marker']}: {scrub(str(error))}") from error + events = self.refresh(fixture) + result = predicate(events) + if result: + return result + if time.monotonic() >= deadline: + markers = {item["marker"] for item in items} + if items and not any(e["marker"] in markers and e["kind"] == "WRAPPER_ENTER" for e in events): + raise CollectionError(f"No runtime-entry evidence for {fixture}; inspect invocations and CloudWatch artifacts") + raise category(f"Evidence deadline exceeded for {fixture}") + time.sleep(1) + + def events_for(self, item): + return selected(list(self.events.values()), marker=item["marker"]) + + def finish(self, item, expected="SUCCEEDED", seconds=100): + try: + result = item["future"].result(timeout=seconds) + except (concurrent.futures.TimeoutError, subprocess.TimeoutExpired) as failure: + raise CollectionError("Client HTTP/driver timeout; not server invocation timeout evidence") from failure + events = self.events_for(item) + arn = result["headers"].get("DurableExecutionArn") + if not arn and events: + arn = events[0]["executionArn"] + if not arn: + raise CollectionError("No durable execution ARN") + item["arn"] = arn + final = aws("lambda", "get-durable-execution", {"DurableExecutionArn": arn}) + save(self.artifacts / "executions" / (item["marker"] + ".json"), final) + if expected: + require(final["Status"] == expected, f"{item['marker']}: durable status {final['Status']}, expected {expected}") + if expected == "SUCCEEDED": + require(json.loads(final["Result"]) == item["marker"], "Execution returned another input's result") + return final + + def history(self, item): + events, marker = [], None + arn = item.get("arn") or self.events_for(item)[0]["executionArn"] + while True: + request = {"DurableExecutionArn": arn, "IncludeExecutionData": True} + if marker: + request["Marker"] = marker + result = aws("lambda", "get-durable-execution-history", request) + events.extend(result.get("Events", [])) + marker = result.get("NextMarker") + if not marker: + break + save(self.artifacts / "histories" / (item["marker"] + ".json"), events) + return events + + def release_all(self): + for gate in list(self.gates): + self.gate(gate, release=True) + self.gates.clear() + + def close(self): + self.release_all() + # All calls have finite HTTP and subprocess budgets; preserve late evidence before teardown. + self.pool.shutdown(wait=True) diff --git a/lmi-tests/pom.xml b/lmi-tests/pom.xml new file mode 100644 index 000000000..75b0ac990 --- /dev/null +++ b/lmi-tests/pom.xml @@ -0,0 +1,40 @@ + + + 4.0.0 + + software.amazon.lambda.durable + aws-durable-execution-sdk-java-parent + 2.2.1-SNAPSHOT + + aws-durable-execution-sdk-java-lmi-tests + LMI cloud regression fixtures + true + + + software.amazon.lambda.durable + aws-durable-execution-sdk-java + ${project.version} + + + org.junit.jupiter + junit-jupiter + test + + + + lmi-fixtures + + + org.apache.maven.plugins + maven-shade-plugin + packageshade + + false + + + + + + + diff --git a/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/FixtureInput.java b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/FixtureInput.java new file mode 100644 index 000000000..a51cea556 --- /dev/null +++ b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/FixtureInput.java @@ -0,0 +1,14 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.lmi; + +/** Test controls are deliberately separate from durable operation identity. */ +public record FixtureInput( + String runId, + String cohort, + String scenario, + String marker, + String controlUrl, + String targetEnvironment, + int peers, + int holdMillis) {} diff --git a/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/InvocationTrace.java b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/InvocationTrace.java new file mode 100644 index 000000000..3a49a9ab4 --- /dev/null +++ b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/InvocationTrace.java @@ -0,0 +1,116 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.lmi; + +import com.amazonaws.services.lambda.runtime.Context; +import java.lang.management.ManagementFactory; +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.UUID; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicLong; +import software.amazon.lambda.durable.serde.JacksonSerDes; + +/** Observes Java stack boundaries, independently of logical future completion. */ +final class InvocationTrace { + static final String ENVIRONMENT = UUID.randomUUID().toString(); + private static final AtomicLong SEQUENCE = new AtomicLong(); + private static final AtomicInteger LIVE_ROOTS = new AtomicInteger(); + private static final AtomicInteger LIVE_TASKS = new AtomicInteger(); + private static final AtomicInteger LIVE_WRAPPERS = new AtomicInteger(); + private static final JacksonSerDes JSON = new JacksonSerDes(); + final FixtureInput input; + final String requestId; + final String executionArn; + final long deadlineNanos; + final AtomicInteger tasks = new AtomicInteger(); + volatile boolean rootExited; + volatile boolean wrapperReturned; + + InvocationTrace(FixtureInput input, Context context, String executionArn) { + this.input = input; + this.requestId = context.getAwsRequestId(); + this.executionArn = executionArn; + this.deadlineNanos = System.nanoTime() + context.getRemainingTimeInMillis() * 1_000_000L; + } + + void wrapperEnter() { + LIVE_WRAPPERS.incrementAndGet(); + event("WRAPPER_ENTER"); + } + + void rootEnter() { + LIVE_ROOTS.incrementAndGet(); + event("ROOT_ENTER"); + } + + void rootExit() { + rootExited = true; + LIVE_ROOTS.decrementAndGet(); + event("ROOT_EXIT"); + } + + void taskEnter(String name) { + tasks.incrementAndGet(); + LIVE_TASKS.incrementAndGet(); + event("TASK_ENTER", Map.of("name", name)); + } + + void taskExit(String name) { + tasks.decrementAndGet(); + LIVE_TASKS.decrementAndGet(); + event("TASK_EXIT", Map.of("name", name)); + } + + void wrapperExit(String status) { + wrapperReturned = true; + LIVE_WRAPPERS.decrementAndGet(); + event("WRAPPER_RETURN", Map.of("status", status)); + } + + void snapshot(ThreadPoolExecutor pool) { + event( + "SNAPSHOT", + Map.of( + "active", + pool.getActiveCount(), + "queued", + pool.getQueue().size(), + "poolSize", + pool.getPoolSize(), + "threads", + ManagementFactory.getThreadMXBean().getThreadCount())); + } + + void event(String kind) { + event(kind, Map.of()); + } + + synchronized void event(String kind, Map details) { + var data = new LinkedHashMap(); + data.put("kind", kind); + data.put("environment", ENVIRONMENT); + data.put("sequence", SEQUENCE.incrementAndGet()); + data.put("nanos", System.nanoTime()); + data.put("epochMillis", System.currentTimeMillis()); + data.put("remainingMillis", (deadlineNanos - System.nanoTime()) / 1_000_000); + data.put("runId", input.runId()); + data.put("deploymentRunId", System.getenv("LMI_TEST_RUN_ID")); + data.put("commit", System.getenv("LMI_COMMIT")); + data.put("cohort", input.cohort()); + data.put("scenario", input.scenario()); + data.put("marker", input.marker()); + data.put("requestId", requestId); + data.put("executionArn", executionArn); + data.put("thread", Thread.currentThread().getName()); + data.put("tasks", tasks.get()); + data.put("rootExited", rootExited); + data.put("wrapperReturned", wrapperReturned); + data.put("liveRoots", LIVE_ROOTS.get()); + data.put("liveTasks", LIVE_TASKS.get()); + data.put("liveWrappers", LIVE_WRAPPERS.get()); + data.putAll(details); + System.out.println("LMI_TEST " + JSON.serialize(data)); + } +} diff --git a/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/LifecycleHandler.java b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/LifecycleHandler.java new file mode 100644 index 000000000..fd5bffa42 --- /dev/null +++ b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/LifecycleHandler.java @@ -0,0 +1,323 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.lmi; + +import com.amazonaws.services.lambda.runtime.Context; +import com.amazonaws.services.lambda.runtime.RequestStreamHandler; +import java.io.IOException; +import java.io.InputStream; +import java.io.OutputStream; +import java.net.URI; +import java.net.http.HttpClient; +import java.net.http.HttpRequest; +import java.net.http.HttpResponse; +import java.nio.charset.StandardCharsets; +import java.time.Duration; +import java.util.List; +import java.util.Map; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.Executors; +import java.util.concurrent.ScheduledExecutorService; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; +import software.amazon.awssdk.services.lambda.model.OperationType; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.DurableContext; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.config.StepConfig; +import software.amazon.lambda.durable.execution.DurableExecutor; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.retry.RetryStrategies; +import software.amazon.lambda.durable.serde.DurableInputOutputSerDes; +import software.amazon.lambda.durable.serde.JacksonSerDes; + +/** Fault fixtures: blocking models user code/cleanup, never a durable delay. */ +public final class LifecycleHandler implements RequestStreamHandler { + private static final DurableInputOutputSerDes WIRE = new DurableInputOutputSerDes(); + private static final JacksonSerDes JSON = new JacksonSerDes(); + private static final HttpClient HTTP = + HttpClient.newBuilder().connectTimeout(Duration.ofSeconds(2)).build(); + private static final ConcurrentHashMap BARRIERS = new ConcurrentHashMap<>(); + private static final ScheduledExecutorService ESCAPES = Executors.newSingleThreadScheduledExecutor(runnable -> { + var thread = new Thread(runnable, "lmi-test-escape"); + thread.setDaemon(true); + return thread; + }); + private static final StepConfig NO_RETRY = + StepConfig.builder().retryStrategy(RetryStrategies.Presets.NO_RETRY).build(); + private final DurableConfig shared = createSharedConfiguration(); + + private static DurableConfig createSharedConfiguration() { + if ("fixed".equals(System.getenv("LMI_EXECUTOR"))) { + return DurableConfig.builder() + .withExecutorService(Executors.newFixedThreadPool(2)) + .build(); + } + return DurableConfig.defaultConfig(); + } + + @Override + public void handleRequest(InputStream source, OutputStream destination, Context runtime) throws IOException { + var input = WIRE.deserialize( + new String(source.readAllBytes(), StandardCharsets.UTF_8), TypeToken.get(DurableExecutionInput.class)); + var trace = new InvocationTrace(readUserInput(input), runtime, input.durableExecutionArn()); + trace.wrapperEnter(); + var status = "THREW"; + try { + var config = DurableConfig.builder() + .withExecutorService(shared.getExecutorService()) + .withDurableExecutionClient(new ObservedClient(shared.getDurableExecutionClient(), trace)) + .build(); + var response = DurableExecutor.execute( + input, + runtime, + TypeToken.get(FixtureInput.class), + (value, context) -> handle(value, context, trace), + config); + destination.write(WIRE.serialize(response).getBytes(StandardCharsets.UTF_8)); + status = response.status().name(); + } finally { + trace.snapshot((ThreadPoolExecutor) shared.getExecutorService()); + trace.wrapperExit(status); + } + } + + private static FixtureInput readUserInput(DurableExecutionInput input) { + var operation = input.initialExecutionState().operations().stream() + .filter(op -> op.type() == OperationType.EXECUTION) + .findFirst() + .orElseThrow(); + return JSON.deserialize(operation.executionDetails().inputPayload(), TypeToken.get(FixtureInput.class)); + } + + String handle(FixtureInput input, DurableContext context, InvocationTrace trace) { + trace.rootEnter(); + try { + if (!context.isReplaying() + && input.targetEnvironment() != null + && !input.targetEnvironment().equals(InvocationTrace.ENVIRONMENT)) { + trace.event("PLACEMENT_MISS"); + return "PLACEMENT_MISS"; + } + return runScenario(input, context, trace); + } finally { + if ("suspend".equals(input.scenario())) { + trace.event("CLEANUP_ENTER"); + boundedPause(2000); + trace.event("CLEANUP_EXIT"); + } + trace.rootExit(); + } + } + + private String runScenario(FixtureInput input, DurableContext context, InvocationTrace trace) { + return switch (input.scenario()) { + case "baseline", "replay", "suspend" -> replay(input, context, trace); + case "hold", "probe" -> context.step("held-step", String.class, step -> hold(input, trace), NO_RETRY); + case "timeout", "failure-inflight", "return-inflight" -> inFlight(input, context, trace); + case "stubborn" -> stubborn(input, context, trace); + case "fixed", "nested" -> fixed(input, context, trace); + case "success" -> context.step("success", String.class, step -> body(trace, "success", input.marker())); + case "failure" -> context.step("failure", String.class, step -> fail(trace, input.marker()), NO_RETRY); + default -> throw new IllegalArgumentException("Unknown fixture scenario"); + }; + } + + private String replay(FixtureInput input, DurableContext context, InvocationTrace trace) { + var value = context.step("success", String.class, step -> body(trace, "success", input.marker())); + try { + context.step("failure", String.class, step -> fail(trace, input.marker()), NO_RETRY); + } catch (IllegalStateException expected) { + trace.event("STORED_FAILURE", Map.of("message", expected.getMessage())); + } + context.wait("resume", Duration.ofSeconds(3)); + return context.step("after-resume", String.class, step -> body(trace, "after-resume", value)); + } + + private static String body(InvocationTrace trace, String name, String value) { + trace.taskEnter(name); + try { + trace.event("BODY", Map.of("name", name, "effectKey", trace.executionArn + "/" + name, "value", value)); + return value; + } finally { + trace.taskExit(name); + } + } + + private static String fail(InvocationTrace trace, String marker) { + body(trace, "failure", marker); + throw new IllegalStateException("expected:" + marker); + } + + private String inFlight(FixtureInput input, DurableContext context, InvocationTrace trace) { + var entered = new CountDownLatch(1); + context.stepAsync("inflight", String.class, step -> blocked(input, trace, entered), NO_RETRY); + await(entered, 5000); + trace.event("ROOT_RESULT"); + if ("failure-inflight".equals(input.scenario())) { + throw new IllegalStateException("expected:" + input.marker()); + } + return input.marker(); + } + + private static String blocked(FixtureInput input, InvocationTrace trace, CountDownLatch entered) { + trace.taskEnter("inflight"); + entered.countDown(); + try { + if (!new CountDownLatch(1).await(Math.min(input.holdMillis(), 120_000), TimeUnit.MILLISECONDS)) { + trace.event( + "timeout".equals(input.scenario()) ? "ESCAPE" : "WORK_COMPLETED", + Map.of("reason", "bounded task completed")); + } + return input.marker(); + } catch (InterruptedException interrupted) { + trace.event("INTERRUPTED"); + Thread.currentThread().interrupt(); + throw new IllegalStateException("fixture interrupted", interrupted); + } finally { + trace.taskExit("inflight"); + } + } + + private String stubborn(FixtureInput input, DurableContext context, InvocationTrace trace) { + var entered = new CountDownLatch(1); + context.runInChildContextAsync("stubborn-child", String.class, child -> { + trace.taskEnter("stubborn-child"); + entered.countDown(); + try { + var end = System.nanoTime() + Math.min(input.holdMillis(), 120_000) * 1_000_000L; + while (System.nanoTime() < end) { + try { + new CountDownLatch(1).await(200, TimeUnit.MILLISECONDS); + } catch (InterruptedException ignored) { + trace.event("IGNORED_INTERRUPT"); + } + } + trace.event("RESIDUAL_EXIT"); + try { + child.step("late-work", String.class, step -> body(trace, "late-work", input.marker())); + trace.event("LATE_ACCEPTED"); + } catch (Throwable rejected) { + trace.event( + "LATE_REJECTED", + Map.of("errorType", rejected.getClass().getSimpleName())); + throw rejected; + } + return input.marker(); + } finally { + trace.taskExit("stubborn-child"); + } + }); + await(entered, 5000); + return input.marker(); + } + + private String fixed(FixtureInput input, DurableContext context, InvocationTrace trace) { + var pool = (ThreadPoolExecutor) shared.getExecutorService(); + var barrier = BARRIERS.computeIfAbsent(input.cohort(), key -> { + ESCAPES.schedule(() -> BARRIERS.remove(key), 20, TimeUnit.SECONDS); + return new CountDownLatch(input.peers()); + }); + trace.event("BARRIER_ENTER"); + barrier.countDown(); + await(barrier, 8000); + trace.event("BARRIER_PASSED"); + var escape = ESCAPES.schedule( + () -> { + trace.snapshot(pool); + trace.event("ESCAPE", Map.of("reason", "fixed executor failed to progress")); + pool.setMaximumPoolSize(32); + pool.setCorePoolSize(32); // Test-only escape after the progress budget. + }, + 8, + TimeUnit.SECONDS); + try { + var value = context.step("success", String.class, step -> body(trace, "success", input.marker())); + if ("nested".equals(input.scenario())) { + nested(context, trace, value); + } + trace.event("PROGRESS"); + return value; + } finally { + escape.cancel(false); + } + } + + private static void nested(DurableContext context, InvocationTrace trace, String value) { + context.runInChildContext("child", String.class, child -> { + trace.taskEnter("child"); + try { + var mapped = child.map( + "map", + List.of(value, value), + String.class, + (item, index, branch) -> + branch.step("mapped-step", String.class, step -> body(trace, "mapped-step", item))); + try (var parallel = child.parallel("parallel")) { + parallel.branch( + "left", String.class, branch -> branch.step("left-step", String.class, step -> value)); + parallel.branch( + "right", String.class, branch -> branch.step("right-step", String.class, step -> value)); + parallel.get(); + } + return mapped.getResult(0); + } finally { + trace.taskExit("child"); + } + }); + } + + private static String hold(FixtureInput input, InvocationTrace trace) { + trace.taskEnter("held-step"); + var end = System.nanoTime() + Math.min(input.holdMillis(), 120_000) * 1_000_000L; + try { + while (System.nanoTime() < end) { + var request = HttpRequest.newBuilder(URI.create(input.controlUrl())) + .timeout(Duration.ofSeconds(2)) + .GET() + .build(); + try { + var response = HTTP.send(request, HttpResponse.BodyHandlers.ofString()); + if (response.statusCode() != 200) { + throw new IllegalStateException("Control object HTTP " + response.statusCode()); + } + if (response.body().trim().equals("release")) { + return input.marker(); + } + } catch (IOException failure) { + throw new IllegalStateException("Control object read failed", failure); + } + trace.event("HEARTBEAT"); + new CountDownLatch(1).await(500, TimeUnit.MILLISECONDS); + } + trace.event("ESCAPE", Map.of("reason", "held step reached escape")); + throw new IllegalStateException("Placement/control deadline expired"); + } catch (InterruptedException interrupted) { + trace.event("INTERRUPTED"); + Thread.currentThread().interrupt(); + throw new IllegalStateException("healthy holder interrupted", interrupted); + } finally { + trace.taskExit("held-step"); + } + } + + private static void await(CountDownLatch latch, long millis) { + try { + if (!latch.await(millis, TimeUnit.MILLISECONDS)) { + throw new IllegalStateException("PLACEMENT_PRECONDITION: barrier not established"); + } + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new IllegalStateException("fixture interrupted", interrupted); + } + } + + private static void boundedPause(long millis) { + try { + new CountDownLatch(1).await(millis, TimeUnit.MILLISECONDS); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + } + } +} diff --git a/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/ObservedClient.java b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/ObservedClient.java new file mode 100644 index 000000000..34c401c98 --- /dev/null +++ b/lmi-tests/src/main/java/software/amazon/lambda/durable/lmi/ObservedClient.java @@ -0,0 +1,47 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.lmi; + +import java.util.List; +import java.util.Map; +import software.amazon.awssdk.services.lambda.model.CheckpointDurableExecutionResponse; +import software.amazon.awssdk.services.lambda.model.GetDurableExecutionStateResponse; +import software.amazon.awssdk.services.lambda.model.OperationUpdate; +import software.amazon.lambda.durable.client.DurableExecutionClient; + +/** Delegates to the real backend without logging checkpoint tokens or payloads. */ +final class ObservedClient implements DurableExecutionClient { + private final DurableExecutionClient delegate; + private final InvocationTrace trace; + + ObservedClient(DurableExecutionClient delegate, InvocationTrace trace) { + this.delegate = delegate; + this.trace = trace; + } + + @Override + public CheckpointDurableExecutionResponse checkpoint(String arn, String token, List updates) { + trace.event( + "CHECKPOINT_CALL", + Map.of( + "operations", + updates.stream() + .map(update -> update.id() + ":" + update.type() + ":" + update.action()) + .toList())); + try { + return delegate.checkpoint(arn, token, updates); + } finally { + trace.event("CHECKPOINT_EXIT"); + } + } + + @Override + public GetDurableExecutionStateResponse getExecutionState(String arn, String token, String marker) { + trace.event("POLL_CALL"); + try { + return delegate.getExecutionState(arn, token, marker); + } finally { + trace.event("POLL_EXIT"); + } + } +} diff --git a/lmi-tests/tests/test_aws_cli.py b/lmi-tests/tests/test_aws_cli.py new file mode 100644 index 000000000..4dd3fa257 --- /dev/null +++ b/lmi-tests/tests/test_aws_cli.py @@ -0,0 +1,84 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Exercise the actual CLI's custom invoke command against a local unsigned endpoint.""" +import json +from http.server import BaseHTTPRequestHandler, HTTPServer +from pathlib import Path +import shutil +import sys +from tempfile import TemporaryDirectory +import threading +import unittest +from unittest.mock import patch +from urllib.parse import unquote + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +import cloud_support +from cloud_support import Cloud + + +@unittest.skipUnless(shutil.which("aws"), "AWS CLI is required for the local invocation contract test") +class AwsCliInvocationTest(unittest.TestCase): + def test_sync_and_async_invocations_send_payload_bytes_and_preserve_results(self): + received = [] + execution_arn = "arn:aws:lambda:us-west-2:123456789012:function:test/durable-execution/test/uuid" + function_arn = "arn:aws:lambda:us-west-2:123456789012:function:test:$LATEST.PUBLISHED" + + class Endpoint(BaseHTTPRequestHandler): + def do_POST(self): + body = self.rfile.read(int(self.headers["Content-Length"])) + payload = json.loads(body) + invocation_type = self.headers["X-Amz-Invocation-Type"] + received.append((unquote(self.path), payload, invocation_type, self.headers.get("Authorization"))) + response = json.dumps(payload["marker"]).encode() if invocation_type == "RequestResponse" else b"" + self.send_response(200 if invocation_type == "RequestResponse" else 202) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(response))) + self.send_header("X-Amz-Durable-Execution-Arn", execution_arn) + self.end_headers() + self.wfile.write(response) + + def log_message(self, *args): + pass + + server = HTTPServer(("127.0.0.1", 0), Endpoint) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + real_aws = cloud_support.aws + + def local_aws(service, operation, data=None, extra=(), **kwargs): + self.assertEqual(("lambda", "invoke"), (service, operation)) + kwargs["timeout"] = 5 + return real_aws(service, operation, data, extra=[*extra, + "--endpoint-url", f"http://127.0.0.1:{server.server_port}", + "--region", "us-west-2", "--no-sign-request"], **kwargs) + + try: + with TemporaryDirectory() as directory, patch("cloud_support.aws", side_effect=local_aws): + cloud = Cloud({"functions": {"default1": {"arn": function_arn}}}, directory) + try: + for scenario, invocation_type in [("baseline", "RequestResponse"), ("timeout", "Event")]: + with self.subTest(scenario=scenario): + payload = {"runId": "test", "scenario": scenario, "marker": scenario + "-λ", + "controlUrl": "https://example.test/control?signature=example"} + result = cloud._invoke("default1", payload) + path, body, actual_type, authorization = received[-1] + self.assertIn(function_arn, path) + self.assertEqual(payload, body) + self.assertEqual(invocation_type, actual_type) + self.assertIsNone(authorization) + self.assertEqual(execution_arn, result["headers"]["DurableExecutionArn"]) + self.assertEqual(payload["marker"] if invocation_type == "RequestResponse" else "", result["body"]) + report = json.loads((Path(directory) / "invocations" / (payload["marker"] + ".json")).read_text()) + self.assertEqual("RETURNED", report["state"]) + self.assertEqual(function_arn, report["functionArn"]) + finally: + cloud.close() + finally: + server.shutdown() + server.server_close() + thread.join(timeout=2) + + +if __name__ == "__main__": + unittest.main() diff --git a/lmi-tests/tests/test_evidence.py b/lmi-tests/tests/test_evidence.py new file mode 100644 index 000000000..43e4960ed --- /dev/null +++ b/lmi-tests/tests/test_evidence.py @@ -0,0 +1,266 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 +import json +from concurrent.futures import Future +from pathlib import Path +import sys +import unittest +from tempfile import TemporaryDirectory + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from cloud_support import (Cloud, CollectionError, PreconditionError, assert_fixed, assert_lifecycle, + assert_overlap, assert_replay, diagnostic, scrub) +from cloud_suite import FIXTURES, cases_for_fixture, verify_function_scaling, run_tests, template +from unittest.mock import Mock, patch + + +def event(kind, sequence, request="a", environment="jvm", **extra): + return {"kind": kind, "sequence": sequence, "nanos": sequence * 1_000_000, + "requestId": request, "marker": request, "environment": environment, + "rootExited": True, "tasks": 0, **extra} + + +class EvidenceTest(unittest.TestCase): + def test_simultaneous_requests_in_different_jvms_are_not_concurrency_evidence(self): + events = [event("TASK_ENTER", 1), event("TASK_ENTER", 1, "b", "other")] + with self.assertRaises(PreconditionError): + assert_overlap(events, {"a", "b"}, 2) + + def test_sequential_requests_in_same_jvm_are_not_overlap(self): + events = [event("TASK_ENTER", 1), event("TASK_EXIT", 2), event("TASK_ENTER", 3, "b")] + with self.assertRaises(PreconditionError): + assert_overlap(events, {"a", "b"}, 2) + + def test_multiple_tasks_from_one_request_do_not_count_as_multiple_invocations(self): + with self.assertRaises(PreconditionError): + assert_overlap([event("TASK_ENTER", 1), event("TASK_ENTER", 2)], {"a"}, 2) + + def test_requires_original_environment_for_recovery(self): + events = [event("TASK_ENTER", 1, "a", "replacement"), event("TASK_ENTER", 2, "b", "replacement")] + with self.assertRaises(PreconditionError): + assert_overlap(events, {"a", "b"}, 2, "original") + + def test_positive_same_jvm_overlap(self): + events = [event("TASK_ENTER", 1), event("TASK_ENTER", 2, "b"), event("TASK_EXIT", 3)] + self.assertEqual("jvm", assert_overlap(events, {"a", "b"}, 2)) + + def test_rejects_pending_before_actual_root_exit(self): + events = [event("WRAPPER_RETURN", 1, status="PENDING", rootExited=False), event("ROOT_EXIT", 2)] + with self.assertRaisesRegex(AssertionError, "before root exit"): + assert_lifecycle(events) + + def test_missing_markers_are_not_a_pass(self): + with self.assertRaises(AssertionError): + assert_lifecycle([]) + with self.assertRaises(AssertionError): + assert_lifecycle([event("WRAPPER_RETURN", 1, status="PENDING")]) + + def test_log_arrival_order_does_not_change_causal_order(self): + assert_lifecycle([event("WRAPPER_RETURN", 2, status="PENDING"), event("ROOT_EXIT", 1)]) + + def test_rejects_late_checkpoints_even_after_final_success(self): + events = [event("ROOT_EXIT", 1), event("WRAPPER_RETURN", 2, status="SUCCEEDED"), event("CHECKPOINT_CALL", 3)] + with self.assertRaisesRegex(AssertionError, "continued after"): + assert_lifecycle(events) + + def test_rejects_live_tasks_at_normal_return(self): + events = [event("ROOT_EXIT", 1), event("WRAPPER_RETURN", 2, status="SUCCEEDED", tasks=1)] + with self.assertRaisesRegex(AssertionError, "live invocation tasks"): + assert_lifecycle(events) + + def test_escape_cannot_make_deadlock_look_successful(self): + events = [event("BARRIER_ENTER", 1), event("BARRIER_ENTER", 2, "b"), + event("BARRIER_PASSED", 3), event("BARRIER_PASSED", 4, "b"), + event("ESCAPE", 5), event("PROGRESS", 6), event("PROGRESS", 7, "b")] + with self.assertRaisesRegex(AssertionError, "starved"): + assert_fixed(events, {"a", "b"}) + + def test_fixed_pool_positive_control(self): + events = [event("BARRIER_ENTER", 1), event("BARRIER_ENTER", 2, "b"), + event("BARRIER_PASSED", 3), event("BARRIER_PASSED", 4, "b"), + event("PROGRESS", 5), event("PROGRESS", 6, "b")] + assert_fixed(events, {"a", "b"}) + + def test_replay_rejects_second_execution_of_checkpointed_body(self): + events, history = self.replay_evidence() + events.append(event("BODY", 9, name="success", value="a")) + with self.assertRaisesRegex(AssertionError, "body ran 2"): + assert_replay(events, history, "a") + + def test_replay_checks_real_history_identity_and_stored_failure(self): + events, history = self.replay_evidence() + assert_replay(events, history, "a") + history.append({"Name": "success", "Id": "changed", "EventType": "StepStarted"}) + with self.assertRaisesRegex(AssertionError, "identity changed"): + assert_replay(events, history, "a") + + def test_replay_cannot_pass_without_resume(self): + events, history = self.replay_evidence() + events = [e for e in events if e["requestId"] != "resume"] + with self.assertRaisesRegex(AssertionError, "No real invocation resume"): + assert_replay(events, history, "a") + + def replay_evidence(self): + events = [event("WRAPPER_ENTER", 1), event("BODY", 2, name="success", value="a"), + event("BODY", 3, name="failure", value="a"), + event("STORED_FAILURE", 4, message="expected:a"), + event("WRAPPER_RETURN", 5, status="PENDING"), + event("WRAPPER_ENTER", 6, "resume", marker="a"), + event("STORED_FAILURE", 7, "resume", marker="a", message="expected:a")] + history = [{"Name": "success", "Id": "1", "EventType": "StepSucceeded"}, + {"Name": "failure", "Id": "2", "EventType": "StepFailed"}, + {"Name": "resume", "Id": "3", "EventType": "WaitSucceeded"}] + return events, history + + def test_structured_lmi_log_envelope(self): + record = event("ROOT_EXIT", 1) + self.assertEqual(record, diagnostic(json.dumps({"message": "LMI_TEST " + json.dumps(record)}))) + self.assertIsNone(diagnostic('{"type":"platform.report"}')) + self.assertIsNone(diagnostic('LMI_TEST malformed')) + + def test_artifacts_redact_control_credentials_in_nested_payloads(self): + value = {"controlUrl": "secret", "InputPayload": '{"controlUrl":"https://bucket/key?X-Amz-Signature=secret"}'} + self.assertNotIn("secret", json.dumps(scrub(value))) + + def test_one_stack_declares_all_five_functions_and_their_scaling_limits(self): + manifest = {"stack": "test", "role": "role", "bucket": "bucket", "provider": "provider", + "commit": "sha", "invocationTimeout": 60, "codeSha256": "digest", "codeKey": "code/digest.jar", "runId": "run-1"} + spec = template(manifest) + functions = [r["Properties"] for r in spec["Resources"].values() if r["Type"] == "AWS::Lambda::Function"] + self.assertEqual(5, len(functions)) + self.assertEqual(set(FIXTURES), set(spec["Outputs"])) + self.assertEqual(5, sum(r["Type"] == "AWS::Logs::LogGroup" for r in spec["Resources"].values())) + concurrencies = set() + for function in functions: + capacity = function["CapacityProviderConfig"]["LambdaManagedInstancesCapacityProviderConfig"] + concurrencies.add(capacity["PerExecutionEnvironmentMaxConcurrency"]) + self.assertEqual("java25", function["Runtime"]) + self.assertEqual(["arm64"], function["Architectures"]) + self.assertEqual(2048, function["MemorySize"]) + self.assertEqual(2, capacity["ExecutionEnvironmentMemoryGiBPerVCpu"]) + self.assertEqual({"MinExecutionEnvironments": 1, "MaxExecutionEnvironments": 1}, + function["FunctionScalingConfig"], "Set limits at creation, before version stabilization") + self.assertEqual(240, function["DurableConfig"]["ExecutionTimeout"]) + self.assertGreater(function["DurableConfig"]["ExecutionTimeout"], function["Timeout"]) + types = {r["Type"] for r in spec["Resources"].values()} + self.assertNotIn("AWS::Lambda::Version", types, "A numbered version duplicates the LMI environment floor") + self.assertNotIn("AWS::Lambda::CapacityProvider", types) + self.assertIn(":$LATEST.PUBLISHED", json.dumps(spec["Outputs"])) + self.assertEqual({1, 2, 8}, concurrencies) + + def test_full_suite_preserves_all_cases_on_their_own_fixture(self): + cloud = Mock() + names = [] + patches = [patch("cloud_suite." + name) for name in + ["replay_case", "overlap_case", "fixed_case", "timeout_case", "inflight_case", "warm_case"]] + mocks = [p.start() for p in patches] + try: + for fixture in FIXTURES: + cases = cases_for_fixture(cloud, fixture) + self.assertTrue(cases) + names.extend(name for name, _ in cases) + for _, run in cases: + run() + calls = [call for mock in mocks for call in mock.call_args_list] + self.assertTrue(all(call.args[0] is cloud and call.args[1] == fixture for call in calls)) + for mock in mocks: + mock.reset_mock() + finally: + for p in patches: + p.stop() + self.assertEqual(13, len(names)) + self.assertEqual(13, len(set(names))) + + @patch("cloud_suite.Cloud") + def test_cloud_suite_requires_all_five_functions(self, cloud): + with TemporaryDirectory() as directory: + path = Path(directory) / "manifest.json" + path.write_text(json.dumps({"functions": {"default1": {}}})) + with patch("cloud_suite.MANIFEST", path), self.assertRaisesRegex(PreconditionError, "All five"): + run_tests() + cloud.assert_not_called() + + @patch("cloud_support.aws", side_effect=RuntimeError("Invoke denied https://example.test/?signature=private-value")) + def test_failed_invocation_is_preserved_in_redacted_artifacts(self, api): + with TemporaryDirectory() as directory: + cloud = Cloud({"functions": {"default1": {"arn": "function"}}}, directory) + try: + with self.assertRaisesRegex(RuntimeError, "Invoke denied"): + cloud._invoke("default1", {"marker": "test", "scenario": "baseline"}) + report = (Path(directory) / "invocations/test.json").read_text() + self.assertIn("Invoke denied", report) + self.assertNotIn("private-value", report) + finally: + cloud.close() + + def test_failed_invocation_is_not_misreported_as_an_sdk_lifecycle_assertion(self): + with TemporaryDirectory() as directory: + cloud = Cloud({}, directory) + cloud.refresh = Mock(return_value=[]) + future = Future() + future.set_exception(RuntimeError("Invoke denied")) + try: + with self.assertRaisesRegex(CollectionError, "Invoke denied"): + cloud.poll("default1", lambda events: False, items=[{"future": future, "marker": "test"}]) + cloud.refresh.assert_not_called() + finally: + cloud.close() + + def test_missing_runtime_entry_is_not_an_sdk_regression(self): + with TemporaryDirectory() as directory: + cloud = Cloud({}, directory) + cloud.refresh = Mock(return_value=[]) + future = Future() + future.set_result({"headers": {"StatusCode": 202}}) + try: + with self.assertRaisesRegex(CollectionError, "No runtime-entry evidence"): + cloud.poll("default1", lambda events: False, seconds=0, + items=[{"future": future, "marker": "test"}]) + cloud.refresh.return_value = [{"kind": "WRAPPER_ENTER", "marker": "test"}] + with self.assertRaisesRegex(AssertionError, "Evidence deadline exceeded"): + cloud.poll("default1", lambda events: False, seconds=0, + items=[{"future": future, "marker": "test"}]) + finally: + cloud.close() + + @patch("cloud_support.aws") + def test_runtime_evidence_must_match_the_current_deployment(self, api): + entry = event("WRAPPER_ENTER", 1, runId="run-2", deploymentRunId="run-1", commit="commit-1") + api.return_value = {"events": [{"eventId": "event", "message": "LMI_TEST " + json.dumps(entry)}]} + with TemporaryDirectory() as directory: + cloud = Cloud({"runId": "run-2", "commit": "commit-2", + "functions": {"default1": {"logGroup": "logs"}}}, directory) + try: + with self.assertRaisesRegex(CollectionError, "outdated deployment"): + cloud.refresh("default1") + self.assertTrue((Path(directory) / "diagnostics.json").exists()) + finally: + cloud.close() + + @patch("cloud_suite.save") + @patch("cloud_suite.time.sleep") + @patch("cloud_suite.aws") + def test_scaling_waits_for_applied_limit_and_an_active_version(self, api, sleep, save): + desired = {"MinExecutionEnvironments": 1, "MaxExecutionEnvironments": 1} + api.side_effect = [{"RequestedFunctionScalingConfig": desired, + "AppliedFunctionScalingConfig": {"MinExecutionEnvironments": 3}}, + {"State": "Active"}, + {"AppliedFunctionScalingConfig": desired}, {"State": "Pending"}, + {"AppliedFunctionScalingConfig": desired}, {"State": "Active"}] + verify_function_scaling("arn:aws:lambda:us-west-2:123456789012:function:test:$LATEST.PUBLISHED", "default2") + self.assertEqual(2, sleep.call_count) + request = api.call_args_list[0].args[2] + self.assertEqual("$LATEST.PUBLISHED", request["Qualifier"]) + self.assertEqual(["get-function-scaling-config", "get-function-configuration"] * 3, + [call.args[1] for call in api.call_args_list]) + + @patch("cloud_suite.save") + @patch("cloud_suite.aws") + def test_failed_version_is_a_setup_failure(self, api, save): + api.side_effect = [{}, {"State": "Failed", "StateReason": "capacity exhausted"}] + with self.assertRaisesRegex(PreconditionError, "capacity exhausted"): + verify_function_scaling("arn:aws:lambda:us-west-2:123456789012:function:test:$LATEST.PUBLISHED", "default2") + + +if __name__ == "__main__": + unittest.main() diff --git a/lmi-tests/tests/test_persistence.py b/lmi-tests/tests/test_persistence.py new file mode 100644 index 000000000..862004be7 --- /dev/null +++ b/lmi-tests/tests/test_persistence.py @@ -0,0 +1,160 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# SPDX-License-Identifier: Apache-2.0 +import json +from pathlib import Path +import sys +from tempfile import TemporaryDirectory +import unittest +from unittest.mock import patch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from cloud_suite import FIXTURES, OWNER, deploy_fixtures, ensure_bucket, template, wait_for_idle_functions +from cloud_support import Cloud, PreconditionError + + +class PersistenceTest(unittest.TestCase): + def manifest(self, run_id="run-1", code_key="code/digest.jar"): + return {"stack": "persistent", "runId": run_id, "codeKey": code_key, "functions": {}, + "account": "123456789012", "region": "us-west-2", "role": "role", "bucket": "bucket", + "provider": "provider", "commit": "sha", "invocationTimeout": 60, "codeSha256": "digest"} + + def stack(self, status="CREATE_COMPLETE"): + return {"StackStatus": status, "Tags": [{"Key": "Suite", "Value": OWNER}], + "Outputs": [{"OutputKey": key, "OutputValue": "function:" + key + ":$LATEST.PUBLISHED"} + for key in FIXTURES]} + + @patch("cloud_suite.save") + @patch("cloud_suite.wait_stack") + @patch("cloud_suite.wait_for_idle_functions") + @patch("cloud_suite.record_fixture") + @patch("cloud_suite.verify_function_scaling") + @patch("cloud_suite.existing_stack") + @patch("cloud_suite.aws") + def test_later_run_updates_same_stack_and_resource_names(self, api, existing, verify, record, idle, wait, save): + existing.side_effect = [None, self.stack(), self.stack(), self.stack("UPDATE_COMPLETE")] + def record_one(manifest, fixture, arn): + manifest["functions"][fixture] = {"arn": arn} + record.side_effect = record_one + deploy_fixtures(self.manifest(), []) + deploy_fixtures(self.manifest("run-2", "code/new-digest.jar"), []) + self.assertEqual(["create-stack", "update-stack"], [call.args[1] for call in api.call_args_list]) + self.assertEqual(["persistent", "persistent"], [call.args[2]["StackName"] for call in api.call_args_list]) + first, second = [json.loads(call.args[2]["TemplateBody"]) for call in api.call_args_list] + self.assertEqual(first["Outputs"], second["Outputs"]) + self.assertEqual(first["Resources"].keys(), second["Resources"].keys()) + for fixture in FIXTURES: + before = first["Resources"][fixture + "Function"]["Properties"] + after = second["Resources"][fixture + "Function"]["Properties"] + self.assertEqual(before["FunctionName"], after["FunctionName"]) + self.assertNotEqual(before["Code"]["S3Key"], after["Code"]["S3Key"]) + self.assertEqual("run-2", after["Environment"]["Variables"]["LMI_TEST_RUN_ID"]) + self.assertNotIn("TimeoutInMinutes", api.call_args_list[1].args[2]) + self.assertEqual(["CREATE_COMPLETE", "UPDATE_COMPLETE"], [call.args[1] for call in wait.call_args_list]) + self.assertEqual(10, verify.call_count) + idle.assert_called_once() + + @patch("cloud_suite.save") + @patch("cloud_suite.wait_stack") + @patch("cloud_suite.wait_for_idle_functions") + @patch("cloud_suite.record_fixture") + @patch("cloud_suite.verify_function_scaling") + @patch("cloud_suite.existing_stack") + @patch("cloud_suite.aws", side_effect=RuntimeError("No updates are to be performed.")) + def test_unchanged_stack_still_gets_readback_without_recreation(self, api, existing, verify, record, idle, wait, save): + existing.return_value = self.stack() + deploy_fixtures(self.manifest(), []) + self.assertEqual("update-stack", api.call_args.args[1]) + wait.assert_not_called() + self.assertEqual(5, verify.call_count) + + @patch("cloud_suite.save") + @patch("cloud_suite.existing_stack") + @patch("cloud_suite.aws") + def test_failed_persistent_stack_is_retained_for_diagnosis(self, api, existing, save): + existing.return_value = self.stack("ROLLBACK_COMPLETE") + with self.assertRaisesRegex(PreconditionError, "retained"): + deploy_fixtures(self.manifest(), []) + api.assert_not_called() + + @patch("cloud_suite.save") + @patch("cloud_suite.existing_stack") + @patch("cloud_suite.aws") + def test_other_stacks_are_not_adopted(self, api, existing, save): + existing.return_value = {"StackStatus": "CREATE_COMPLETE", "Tags": []} + with self.assertRaisesRegex(PreconditionError, "not owned"): + deploy_fixtures(self.manifest(), []) + api.assert_not_called() + + @patch("cloud_suite.save") + @patch("cloud_suite.existing_stack", return_value=None) + @patch("cloud_suite.aws", side_effect=RuntimeError("create failed")) + def test_failed_creation_does_not_trigger_deletion(self, api, existing, save): + with self.assertRaisesRegex(RuntimeError, "create failed"): + deploy_fixtures(self.manifest(), []) + self.assertEqual(["create-stack"], [call.args[1] for call in api.call_args_list]) + + @patch("cloud_suite.aws") + def test_expired_budget_prevents_changes(self, api): + with self.assertRaisesRegex(PreconditionError, "budget exhausted"): + deploy_fixtures(self.manifest(), [], seconds=0) + api.assert_not_called() + + @patch("cloud_suite.save") + @patch("cloud_suite.time.sleep") + @patch("cloud_suite.aws") + def test_previous_executions_are_allowed_to_finish_before_update(self, api, sleep, save): + stack = {"Outputs": [self.stack()["Outputs"][0]]} + api.side_effect = [{"DurableExecutions": [{"DurableExecutionArn": "old", "Status": "RUNNING"}]}, + {"DurableExecutions": []}] + wait_for_idle_functions(stack) + self.assertEqual(1, sleep.call_count) + self.assertTrue(all(call.args[1] == "list-durable-executions-by-function" for call in api.call_args_list)) + self.assertEqual(["RUNNING"], api.call_args.args[2]["Statuses"]) + + @patch("cloud_suite.save") + @patch("cloud_suite.aws") + def test_still_running_execution_prevents_update_without_stopping_it(self, api, save): + api.return_value = {"DurableExecutions": [{"DurableExecutionArn": "old", "Status": "RUNNING"}]} + with self.assertRaisesRegex(PreconditionError, "code was not updated"): + wait_for_idle_functions({"Outputs": [self.stack()["Outputs"][0]]}, seconds=0) + self.assertTrue(all(call.args[1].startswith("list-") for call in api.call_args_list)) + + @patch("cloud_suite.save") + @patch("cloud_suite.aws") + def test_incomplete_execution_listing_cannot_be_treated_as_idle(self, api, save): + api.return_value = {"DurableExecutions": [], "NextMarker": "another-page"} + with self.assertRaisesRegex(PreconditionError, "finish checking"): + wait_for_idle_functions({"Outputs": [self.stack()["Outputs"][0]]}, seconds=0) + + @patch("cloud_suite.aws") + def test_owned_bucket_is_reused_and_only_control_objects_expire(self, api): + api.side_effect = [{}, {"TagSet": [{"Key": "Suite", "Value": OWNER}, {"Key": "Stack", "Value": "persistent"}]}, {}, {}] + ensure_bucket(self.manifest(), []) + operations = [call.args[1] for call in api.call_args_list] + self.assertNotIn("create-bucket", operations) + self.assertNotIn("delete-bucket", operations) + lifecycle = api.call_args.args[2]["LifecycleConfiguration"]["Rules"] + self.assertEqual([{"Prefix": "control/"}], [rule["Filter"] for rule in lifecycle]) + + @patch("cloud_suite.aws") + def test_absent_bucket_is_created_once(self, api): + api.side_effect = [RuntimeError("404 Not Found"), {}, {}, {}, {}] + ensure_bucket(self.manifest(), []) + self.assertEqual(1, sum(call.args[1] == "create-bucket" for call in api.call_args_list)) + + @patch("cloud_suite.aws", side_effect=RuntimeError("Access denied")) + def test_bucket_access_error_does_not_trigger_creation(self, api): + with self.assertRaisesRegex(RuntimeError, "Access denied"): + ensure_bucket(self.manifest(), []) + self.assertEqual(1, api.call_count) + + @patch("cloud_support.aws", return_value="https://example.test/control") + def test_control_objects_are_scoped_to_the_run(self, api): + with TemporaryDirectory() as directory: + cloud = Cloud({"runId": "run-2", "bucket": "bucket"}, directory) + try: + cloud.gate("barrier") + self.assertEqual("control/run-2/barrier", api.call_args_list[0].args[2]["Key"]) + self.assertIn("s3://bucket/control/run-2/barrier", api.call_args_list[1].kwargs["extra"]) + finally: + cloud.close() diff --git a/pom.xml b/pom.xml index bb78588a5..be00a7bd5 100644 --- a/pom.xml +++ b/pom.xml @@ -46,6 +46,7 @@ insight-plugin examples conformance-tests + lmi-tests coverage-report diff --git a/sdk/src/main/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionException.java b/sdk/src/main/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionException.java index acd2fd9fc..a28b3f582 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionException.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionException.java @@ -9,14 +9,25 @@ public class UnrecoverableDurableExecutionException extends DurableExecutionExce private final ErrorObject errorObject; private final boolean retryable; - public UnrecoverableDurableExecutionException(ErrorObject errorObject, boolean retryable) { - super(errorObject.errorMessage()); + /** + * Creates an unrecoverable execution exception with an optional retry path and original cause. + * + * @param errorObject serialized error details + * @param retryable whether the current invocation may be retried + * @param cause original failure that caused this exception + */ + public UnrecoverableDurableExecutionException(ErrorObject errorObject, boolean retryable, Throwable cause) { + super(errorObject.errorMessage(), cause); this.errorObject = errorObject; this.retryable = retryable; } + public UnrecoverableDurableExecutionException(ErrorObject errorObject, boolean retryable) { + this(errorObject, retryable, null); + } + public UnrecoverableDurableExecutionException(ErrorObject errorObject) { - this(errorObject, false); + this(errorObject, false, null); } /** Returns the error details for this unrecoverable exception. */ diff --git a/sdk/src/main/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcher.java b/sdk/src/main/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcher.java index 3b0def81f..4a7b08e9f 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcher.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcher.java @@ -8,7 +8,9 @@ import java.util.Objects; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentLinkedQueue; +import java.util.concurrent.ExecutionException; import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; import java.util.function.Consumer; import java.util.function.Function; @@ -23,7 +25,7 @@ * @param Request type */ public class ApiRequestDelayedBatcher { - private static final Duration MAX_DELAY = Duration.ofMinutes(60); + private static final long NO_FLUSH_DEADLINE = Long.MAX_VALUE; /** Maximum items allowed in a single batch */ private final int maxItemCount; @@ -102,8 +104,8 @@ CompletableFuture submit(T request, Duration flushDelay) { } } - /** Flushes pending batch and waits for completion */ - void shutdown() { + /** Flushes pending batches and waits no longer than the supplied timeout. */ + void shutdown(Duration timeout) throws InterruptedException, ExecutionException, TimeoutException { synchronized (delayedBatch) { // cancel the flush timer if it has not been triggered this.delayedBatchFlushTimer.cancel(false); @@ -112,14 +114,15 @@ void shutdown() { } // wait for previous batches to be flushed - flushingQueueFuture.join(); + flushingQueueFuture.get(Math.max(0, timeout.toNanos()), TimeUnit.NANOSECONDS); } /** clear the current batch and creates a new batch */ private void initializeDelayedBatch() { this.delayedBatch.clear(); - // MAX_DELAY is longer than a single Lambda invocation - this.delayedBatchFlushTime = System.nanoTime() + MAX_DELAY.toNanos(); + // No timer is scheduled until the first item supplies an actual flush deadline. Using an unbounded sentinel + // avoids coupling batching behavior to the maximum Lambda invocation duration. + this.delayedBatchFlushTime = NO_FLUSH_DEADLINE; // the timer future is created initially without a timeout until an item is added to the batch this.delayedBatchFlushTimer = new CompletableFuture<>(); diff --git a/sdk/src/main/java/software/amazon/lambda/durable/execution/CheckpointManager.java b/sdk/src/main/java/software/amazon/lambda/durable/execution/CheckpointManager.java index 34f44134c..d5a65c3be 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/execution/CheckpointManager.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/execution/CheckpointManager.java @@ -11,6 +11,8 @@ import java.util.Objects; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.TimeoutException; import java.util.function.BooleanSupplier; import java.util.function.Consumer; import org.slf4j.Logger; @@ -148,8 +150,8 @@ private CompletableFuture pollForUpdateInternal( }); } - /** Cancels all polling futures and waits for all pending checkpoint requests to complete */ - void shutdown() { + /** Cancels polling futures and bounds the wait for pending checkpoint requests. */ + void shutdown(Duration timeout) throws InterruptedException, ExecutionException, TimeoutException { // complete all polling futures with an exception List>> allFutures; synchronized (pollingFutures) { @@ -162,7 +164,7 @@ void shutdown() { } // wait for all non-polling checkpoint requests to complete - checkpointApiRequestDelayedBatcher.shutdown(); + checkpointApiRequestDelayedBatcher.shutdown(timeout); } /** diff --git a/sdk/src/main/java/software/amazon/lambda/durable/execution/DurableExecutor.java b/sdk/src/main/java/software/amazon/lambda/durable/execution/DurableExecutor.java index d8db91326..7536855d7 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/execution/DurableExecutor.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/execution/DurableExecutor.java @@ -5,8 +5,6 @@ import com.amazonaws.services.lambda.runtime.Context; import com.amazonaws.services.lambda.runtime.RequestHandler; import java.nio.charset.StandardCharsets; -import java.util.concurrent.CompletableFuture; -import java.util.concurrent.CompletionException; import java.util.concurrent.atomic.AtomicReference; import java.util.function.BiFunction; import org.slf4j.Logger; @@ -64,138 +62,134 @@ public static DurableExecutionOutput execute( executionManager.registerActiveThread(null); // Captured for onInvocationEnd, which runs outside the handler thread below. var pluginExecutionInput = new AtomicReference<>(); - var handlerFuture = CompletableFuture.supplyAsync( - () -> { - executionManager.setCurrentThreadContext(new ThreadContext(null, ThreadType.CONTEXT)); + var handlerFuture = executionManager.submitRootTask(() -> { + executionManager.setCurrentThreadContext(new ThreadContext(null, ThreadType.CONTEXT)); - // Deserialize once and share the value with the plugin hooks and the handler below. A second - // deserialization would double the cost, hand plugins a different object than the handler, and - // re-run any side effects in a stateful custom SerDes. A failure is captured rather than thrown - // so onInvocationStart still fires before it surfaces, keeping the start/end hooks paired. - // SerDes is a public extension point whose deserialize declares no checked exceptions, so an - // implementation may sneaky-throw one; capture every Throwable and rethrow it unchanged. - I userInput = null; - Throwable inputFailure = null; - try { - userInput = extractUserInput( - executionManager.getExecutionOperation(), config.getSerDes(), inputType); - } catch (Throwable t) { - inputFailure = t; - } - pluginExecutionInput.set(userInput); + // Deserialize once and share the value with the plugin hooks and the handler below. A second + // deserialization would double the cost, hand plugins a different object than the handler, and + // re-run any side effects in a stateful custom SerDes. A failure is captured rather than thrown + // so onInvocationStart still fires before it surfaces, keeping the start/end hooks paired. + // SerDes is a public extension point whose deserialize declares no checked exceptions, so an + // implementation may sneaky-throw one; capture every Throwable and rethrow it unchanged. + I userInput = null; + Throwable inputFailure = null; + try { + userInput = + extractUserInput(executionManager.getExecutionOperation(), config.getSerDes(), inputType); + } catch (Throwable t) { + inputFailure = t; + } + pluginExecutionInput.set(userInput); - // onInvocationStart runs on the user thread so plugins can - // inject ThreadLocal objects, update MDC, etc. - // executionStartTime comes from the initial EXECUTION operation in the first backend event. - if (!pluginRunner.isEmpty()) { - pluginRunner.onInvocationStart(new InvocationInfo( - requestId, - executionArn, - isFirstInvocation, - executionManager.getExecutionOperation().startTimestamp(), - userInput, - PluginInfoConverter.toOperationItemMap( - executionManager.getOperationsSnapshot(), - executionManager.getInitialOperationIds()), - PluginInfoConverter.toOperationItemMap( - executionManager.getUpdatedOperationsSnapshot(), - executionManager.getInitialOperationIds()))); - } - if (inputFailure != null) { - ExceptionHelper.sneakyThrow(inputFailure); - } + // onInvocationStart runs on the user thread so plugins can + // inject ThreadLocal objects, update MDC, etc. + // executionStartTime comes from the initial EXECUTION operation in the first backend event. + if (!pluginRunner.isEmpty()) { + pluginRunner.onInvocationStart(new InvocationInfo( + requestId, + executionArn, + isFirstInvocation, + executionManager.getExecutionOperation().startTimestamp(), + userInput, + PluginInfoConverter.toOperationItemMap( + executionManager.getOperationsSnapshot(), + executionManager.getInitialOperationIds()), + PluginInfoConverter.toOperationItemMap( + executionManager.getUpdatedOperationsSnapshot(), + executionManager.getInitialOperationIds()))); + } + if (inputFailure != null) { + ExceptionHelper.sneakyThrow(inputFailure); + } - var context = DurableContextImpl.createRootContext(executionManager, config, lambdaContext); - DurableContextImpl.setCurrentContext(context); - // use a try-with-resources to clear logger properties - try (var ignored = DurableLogger.attachContext()) { - return handler.apply(userInput, context); - } - }, - config.getExecutorService()); // Get executor from config for running user code + var context = DurableContextImpl.createRootContext(executionManager, config, lambdaContext); + DurableContextImpl.setCurrentContext(context); + // use a try-with-resources to clear logger properties + try (var ignored = DurableLogger.attachContext()) { + return handler.apply(userInput, context); + } + }); - // Execute the handlerFuture in ExecutionManager. If it completes successfully, the output of user function - // will be returned. Otherwise, it will complete exceptionally with a SuspendExecutionException or a - // failure. - try { - return executionManager - .runUntilCompleteOrSuspend(handlerFuture) - .handle((result, ex) -> { - if (ex != null) { - // an exception thrown from handlerFuture or suspension/termination occurred - Throwable cause = ExceptionHelper.unwrapCompletableFuture(ex); + var outcome = executionManager.awaitInvocationOutcome(handlerFuture); + executionManager.beginDraining(); - // return PENDING if it's SuspendExecutionException - if (cause instanceof SuspendExecutionException) { - fireOnInvocationEnd( - pluginRunner, - executionManager, - requestId, - executionArn, - isFirstInvocation, - InvocationStatus.PENDING, - null, - pluginExecutionInput.get(), - null); - return DurableExecutionOutput.pending(); - } - - // let the backend retry the invocation if the exception is retryable - if (cause - instanceof - UnrecoverableDurableExecutionException - unrecoverableDurableExecutionException - && unrecoverableDurableExecutionException.isRetryable()) { - fireOnInvocationEnd( - pluginRunner, - executionManager, - requestId, - executionArn, - isFirstInvocation, - InvocationStatus.RETRYING, - cause, - pluginExecutionInput.get(), - null); - throw unrecoverableDurableExecutionException; - } + String outputPayload = null; + Throwable responseFailure = null; + if (outcome.failure() == null) { + try { + outputPayload = handleLargePayload( + executionManager, config.getSerDes().serialize(outcome.result())); + } catch (Throwable failure) { + responseFailure = ExceptionHelper.unwrapCompletableFuture(failure); + } + } - // fail the execution otherwise - logger.debug("Execution failed: {}", cause.getMessage()); - fireOnInvocationEnd( - pluginRunner, - executionManager, - requestId, - executionArn, - isFirstInvocation, - InvocationStatus.FAILED, - cause, - pluginExecutionInput.get(), - null); - return DurableExecutionOutput.failure(buildErrorObject(cause, config.getSerDes())); - } - // user handler complete successfully - logger.debug("Execution completed"); - var outputPayload = config.getSerDes().serialize(result); - var output = - DurableExecutionOutput.success(handleLargePayload(executionManager, outputPayload)); - fireOnInvocationEnd( - pluginRunner, - executionManager, - requestId, - executionArn, - isFirstInvocation, - InvocationStatus.SUCCEEDED, - null, - pluginExecutionInput.get(), - result); - return output; - }) - .join(); - } catch (CompletionException e) { - // unwrap the CompletionException and rethrow the wrapped exception - ExceptionHelper.sneakyThrow(ExceptionHelper.unwrapCompletableFuture(e)); + var originalFailure = outcome.failure() != null ? outcome.failure() : responseFailure; + var finalFailure = executionManager.finishInvocation(originalFailure); + if (responseFailure != null + && finalFailure == responseFailure + && !(responseFailure instanceof UnrecoverableDurableExecutionException failure + && failure.isRetryable())) { + ExceptionHelper.sneakyThrow(responseFailure); return null; } + + if (finalFailure instanceof SuspendExecutionException) { + fireOnInvocationEnd( + pluginRunner, + executionManager, + requestId, + executionArn, + isFirstInvocation, + InvocationStatus.PENDING, + null, + pluginExecutionInput.get(), + null); + return DurableExecutionOutput.pending(); + } + + if (finalFailure instanceof UnrecoverableDurableExecutionException failure && failure.isRetryable()) { + fireOnInvocationEnd( + pluginRunner, + executionManager, + requestId, + executionArn, + isFirstInvocation, + InvocationStatus.RETRYING, + failure, + pluginExecutionInput.get(), + null); + throw failure; + } + + if (finalFailure != null) { + logger.debug("Execution failed: {}", finalFailure.getMessage()); + fireOnInvocationEnd( + pluginRunner, + executionManager, + requestId, + executionArn, + isFirstInvocation, + InvocationStatus.FAILED, + finalFailure, + pluginExecutionInput.get(), + null); + return DurableExecutionOutput.failure(buildErrorObject(finalFailure, config.getSerDes())); + } + + logger.debug("Execution completed"); + var output = DurableExecutionOutput.success(outputPayload); + fireOnInvocationEnd( + pluginRunner, + executionManager, + requestId, + executionArn, + isFirstInvocation, + InvocationStatus.SUCCEEDED, + null, + pluginExecutionInput.get(), + outcome.result()); + return output; } } @@ -236,14 +230,12 @@ private static String handleLargePayload(ExecutionManager executionManager, Stri LAMBDA_RESPONSE_SIZE_LIMIT); // Checkpoint the large result and wait for it to complete - executionManager - .sendOperationUpdate(OperationUpdate.builder() - .type(OperationType.EXECUTION) - .id(executionManager.getExecutionOperation().id()) - .action(OperationAction.SUCCEED) - .payload(outputPayload) - .build()) - .join(); + executionManager.awaitCheckpointCompletion(executionManager.sendOperationUpdate(OperationUpdate.builder() + .type(OperationType.EXECUTION) + .id(executionManager.getExecutionOperation().id()) + .action(OperationAction.SUCCEED) + .payload(outputPayload) + .build())); // Return empty result, we checkpointed the data manually logger.debug("Execution completed (large response checkpointed)"); diff --git a/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutionManager.java b/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutionManager.java index 0e9d8426e..0112c4b1a 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutionManager.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutionManager.java @@ -3,23 +3,32 @@ package software.amazon.lambda.durable.execution; import com.amazonaws.services.lambda.runtime.Context; +import java.time.Duration; import java.time.Instant; import java.util.ArrayList; import java.util.Collection; import java.util.Collections; +import java.util.Comparator; import java.util.HashSet; import java.util.List; import java.util.Map; import java.util.Objects; +import java.util.Optional; import java.util.Set; -import java.util.concurrent.CancellationException; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; -import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.ExecutionException; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.RejectedExecutionException; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicLong; import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Supplier; import java.util.stream.Collectors; import org.slf4j.Logger; import org.slf4j.LoggerFactory; +import software.amazon.awssdk.services.lambda.model.ErrorObject; import software.amazon.awssdk.services.lambda.model.Operation; import software.amazon.awssdk.services.lambda.model.OperationStatus; import software.amazon.awssdk.services.lambda.model.OperationType; @@ -30,6 +39,7 @@ import software.amazon.lambda.durable.model.SafeCloseable; import software.amazon.lambda.durable.operation.BaseDurableOperation; import software.amazon.lambda.durable.plugin.PluginInfoConverter; +import software.amazon.lambda.durable.util.ExceptionHelper; /** * Central manager for durable execution coordination. @@ -38,7 +48,8 @@ * *
    *
  • Execution state (operations, checkpoint token) - *
  • Thread lifecycle (registration/deregistration) + *
  • Logical thread and executor-task lifecycle + *
  • Invocation deadline and work admission *
  • Checkpoint batching (via CheckpointManager) *
  • Checkpoint result handling (CheckpointManager callback) *
  • Polling (for waits and retries) @@ -55,6 +66,34 @@ public class ExecutionManager implements SafeCloseable { private static final Logger logger = LoggerFactory.getLogger(ExecutionManager.class); + private static final AtomicLong LOCAL_INVOCATION_SEQUENCE = new AtomicLong(); + // Stop waiting for the root early enough to cancel tasks and flush checkpoints before the platform deadline. + private static final Duration EXECUTION_DEADLINE_HEADROOM = Duration.ofMillis(500); + // Leave a final margin for response serialization and return through the Lambda runtime wrapper. + private static final Duration RESPONSE_HEADROOM = Duration.ofMillis(100); + // Lambda always supplies a deadline; this cap also keeps local/null-context cleanup bounded. + private static final Duration MAX_CLEANUP_DURATION = Duration.ofSeconds(30); + private static final Duration MAX_GRACEFUL_DRAIN_DURATION = Duration.ofSeconds(5); + private static final Duration MAX_CANCELLATION_DRAIN_DURATION = Duration.ofSeconds(1); + private static final Duration CHECKPOINT_CLEANUP_RESERVE = Duration.ofMillis(100); + private static final String LIFECYCLE_ERROR_TYPE = + "software.amazon.lambda.durable.execution.InvocationLifecycleException"; + + enum LifecycleState { + /** The invocation accepts new durable operations, executor tasks, checkpoints, and polls. */ + OPEN, + + /** + * Invocation shutdown has started. New operations and tasks are rejected, while already-admitted work may still + * checkpoint or poll as it finishes. + */ + DRAINING, + + /** Invocation cleanup has ended; all new operations, tasks, checkpoints, and polls are rejected. */ + CLOSED + } + + record InvocationOutcome(T result, Throwable failure) {} // ===== Execution State ===== private final Map operationStorage; @@ -66,6 +105,16 @@ public class ExecutionManager implements SafeCloseable { private final Set updatedOperationIdsSinceLastInvocation; private final Set initialOperationIds; + // ===== Invocation Lifecycle ===== + private final Object admissionLock = new Object(); + private final String invocationId; + private final Long deadlineNanos; + private final AtomicLong taskSequence = new AtomicLong(); + private final Map> activeExecutorTasks = new ConcurrentHashMap<>(); + private volatile LifecycleState lifecycleState = LifecycleState.OPEN; + private boolean checkpointAdmissionOpen = true; + private boolean restoreInvocationThreadInterrupt; + // ===== Thread Coordination ===== private final Map registeredOperations = new ConcurrentHashMap<>(); private final Set activeThreads = Collections.synchronizedSet(new HashSet<>()); @@ -81,6 +130,8 @@ public ExecutionManager(DurableExecutionInput input, DurableConfig config, Conte durableConfig = config; this.durableExecutionArn = input.durableExecutionArn(); this.lambdaContext = lambdaContext; + this.invocationId = resolveInvocationId(lambdaContext, durableExecutionArn); + this.deadlineNanos = resolveDeadlineNanos(lambdaContext); // Store the set of operation IDs updated since the last successful invocation this.updatedOperationIdsSinceLastInvocation = @@ -189,7 +240,71 @@ public Collection getUpdatedOperationsSnapshot() { /** Registers an operation so it can receive checkpoint completion notifications. */ public void registerOperation(BaseDurableOperation operation) { - registeredOperations.put(operation.getOperationId(), operation); + synchronized (admissionLock) { + requireOpen("durable operation"); + registeredOperations.put(operation.getOperationId(), operation); + } + } + + /** Submits the invocation's root handler and records its executor task separately from its logical result. */ + CompletableFuture submitRootTask(Supplier action) { + return submitExecutorTask(ExecutorTaskHandle.Role.ROOT, null, durableConfig.getExecutorService(), action); + } + + /** Submits an operation handler and records its executor task separately from its logical result. */ + public CompletableFuture submitOperationTask(BaseDurableOperation operation, Runnable action) { + var role = operation.getType() == OperationType.STEP + ? ExecutorTaskHandle.Role.STEP + : switch (operation.getSubType()) { + case MAP, PARALLEL -> ExecutorTaskHandle.Role.COORDINATOR; + default -> ExecutorTaskHandle.Role.CHILD_CONTEXT; + }; + return submitExecutorTask(role, operation.getOperationId(), durableConfig.getExecutorService(), () -> { + action.run(); + return null; + }); + } + + private CompletableFuture submitExecutorTask( + ExecutorTaskHandle.Role role, String operationId, ExecutorService executor, Supplier action) { + ExecutorTaskHandle task; + synchronized (admissionLock) { + requireOpen("task"); + var taskId = taskSequence.incrementAndGet(); + task = new ExecutorTaskHandle<>(taskId, role, operationId, () -> activeExecutorTasks.remove(taskId)); + activeExecutorTasks.put(task.id(), task); + } + + try { + task.bindExecution(executor.submit(() -> task.run(action))); + return task.completion(); + } catch (RuntimeException | Error failure) { + task.submissionFailed(failure); + throw failure; + } + } + + String getInvocationId() { + return invocationId; + } + + Optional getRemainingInvocationTime() { + if (deadlineNanos == null) { + return Optional.empty(); + } + return Optional.of(Duration.ofNanos(Math.max(0, deadlineNanos - System.nanoTime()))); + } + + LifecycleState getLifecycleState() { + synchronized (admissionLock) { + return lifecycleState; + } + } + + List> getActiveExecutorTasks() { + return activeExecutorTasks.values().stream() + .sorted(Comparator.comparingLong(ExecutorTaskHandle::id)) + .toList(); } // ===== Checkpoint Completion Handler ===== @@ -336,7 +451,7 @@ public void deregisterActiveThread(String threadId) { boolean tryStartCheckpointProcessing() { synchronized (activeThreads) { - if (executionExceptionFuture.isDone()) { + if (executionExceptionFuture.isDone() || lifecycleState == LifecycleState.CLOSED) { return false; } checkpointRequestsInFlight++; @@ -380,7 +495,25 @@ private void preSuspendCheck() { // This method will checkpoint the operation updates to the durable backend and return a future which completes // when the checkpoint completes. public CompletableFuture sendOperationUpdate(OperationUpdate update) { - return checkpointManager.checkpoint(update); + return admitCheckpoint(() -> checkpointManager.checkpoint(update)); + } + + /** Waits for a response-critical checkpoint without crossing the invocation response deadline. */ + T awaitCheckpointCompletion(CompletableFuture checkpointFuture) { + var waitNanos = deadlineNanos == null + ? MAX_CLEANUP_DURATION.toNanos() + : remainingNanos(deadlineNanos - RESPONSE_HEADROOM.toNanos()); + try { + return checkpointFuture.get(waitNanos, TimeUnit.NANOSECONDS); + } catch (InterruptedException interrupted) { + restoreInvocationThreadInterrupt = true; + throw lifecycleFailure("Interrupted while waiting for checkpoint completion", interrupted); + } catch (TimeoutException timeout) { + throw lifecycleFailure("Checkpoint completion exceeded the invocation deadline", timeout); + } catch (ExecutionException failure) { + ExceptionHelper.sneakyThrow(ExceptionHelper.unwrapCompletableFuture(failure.getCause())); + return null; + } } // ===== Polling ===== @@ -391,7 +524,7 @@ public CompletableFuture sendOperationUpdate(OperationUpdate update) { // wait while another thread is still running, and we therefore are not // re-invoked because we never suspended. public CompletableFuture pollForOperationUpdates(String operationId) { - return checkpointManager.pollForUpdate(operationId); + return admitCheckpoint(() -> checkpointManager.pollForUpdate(operationId)); } /** @@ -402,48 +535,169 @@ public CompletableFuture pollForOperationUpdates(String operationId) * @return a completable future that completes with the operation update */ public CompletableFuture pollForOperationUpdates(String operationId, Instant at) { - return checkpointManager.pollForUpdate(operationId, at); + return admitCheckpoint(() -> checkpointManager.pollForUpdate(operationId, at)); } - // ===== Utilities ===== - /** Shutdown the checkpoint batcher. */ + // ===== Invocation Cleanup ===== @Override public void close() { - validateRunningThreads(); - - checkpointManager.shutdown(); - } - - private void validateRunningThreads() { - // This will detect stuck user thread and thread leaks in the thread pool - for (BaseDurableOperation op : registeredOperations.values()) { - var userHandlerFuture = op.getRunningUserHandler(); - if (userHandlerFuture != null && !userHandlerFuture.isDone()) { - // Some user threads can still be running because - // the operations that run them have never been waiting for and the execution has completed. - logger.info("Waiting for operation to complete before shutting down: {}", op.getOperationId()); - try { - userHandlerFuture.get(); - } catch (InterruptedException | CancellationException e) { - // if the user handler is stuck - throw new IllegalStateException( - "Stuck running user handler when shutting down: " + op.getOperationId()); - } catch (Exception e) { - // ok if the future completed exceptionally + var failure = finishInvocation(null); + if (failure != null) { + ExceptionHelper.sneakyThrow(failure); + } + } + + /** Drains invocation-owned tasks and shuts down checkpoint coordination before a response is returned. */ + Throwable finishInvocation(Throwable originalFailure) { + if (lifecycleState == LifecycleState.CLOSED) { + return originalFailure; + } + beginDraining(); + var hardDeadline = cleanupDeadlineNanos(); + Throwable cleanupFailure = null; + try { + cleanupFailure = drainExecutorTasks(originalFailure, hardDeadline); + } catch (InterruptedException interrupted) { + restoreInvocationThreadInterrupt = true; + cleanupFailure = lifecycleFailure("Interrupted while draining invocation tasks", interrupted); + signalLifecycleFailure(cleanupFailure); + cancelActiveTasks(); + } finally { + try { + stopCheckpointAdmission(); + cleanupFailure = shutdownCheckpoints(hardDeadline, cleanupFailure); + } finally { + closeLifecycle(); + if (restoreInvocationThreadInterrupt) { + Thread.currentThread().interrupt(); } } } - // double check if the thread pool is empty - if (durableConfig.getExecutorService() instanceof ThreadPoolExecutor threadPoolExecutor) { - var threadCount = threadPoolExecutor.getActiveCount(); - // This may or may not be a problem because getActiveCount doesn't return an accurate number - if (threadCount > 0) { - logger.warn("{} active threads in user executor pool when shutting down", threadCount); + if (cleanupFailure != null) { + if (originalFailure != null && originalFailure != cleanupFailure) { + cleanupFailure.addSuppressed(originalFailure); + } + return cleanupFailure; + } + return originalFailure; + } + + private Throwable shutdownCheckpoints(long hardDeadline, Throwable priorFailure) { + try { + checkpointManager.shutdown(remainingDuration(hardDeadline)); + return priorFailure; + } catch (InterruptedException interrupted) { + restoreInvocationThreadInterrupt = true; + return combineFailures( + priorFailure, lifecycleFailure("Interrupted while shutting down checkpoints", interrupted)); + } catch (ExecutionException | TimeoutException failure) { + return combineFailures( + priorFailure, lifecycleFailure("Checkpoint shutdown exceeded its cleanup budget", failure)); + } catch (RuntimeException failure) { + return combineFailures(priorFailure, lifecycleFailure("Checkpoint shutdown failed", failure)); + } + } + + private Throwable drainExecutorTasks(Throwable originalFailure, long hardDeadline) throws InterruptedException { + var taskDeadline = Math.max(System.nanoTime(), hardDeadline - CHECKPOINT_CLEANUP_RESERVE.toNanos()); + var retrying = + originalFailure instanceof UnrecoverableDurableExecutionException failure && failure.isRetryable(); + if (!retrying) { + var gracefulDeadline = minDeadline(taskDeadline, MAX_GRACEFUL_DRAIN_DURATION); + if (awaitTasksUntil(gracefulDeadline)) { + return null; + } + + var failure = lifecycleFailure("Invocation tasks did not exit during graceful cleanup", null); + signalLifecycleFailure(failure); + cancelActiveTasks(); + awaitTasksUntil(minDeadline(taskDeadline, MAX_CANCELLATION_DRAIN_DURATION)); + return failure; + } + + signalLifecycleFailure(originalFailure); + cancelActiveTasks(); + if (awaitTasksUntil(minDeadline(taskDeadline, MAX_CANCELLATION_DRAIN_DURATION))) { + return null; + } + return lifecycleFailure("Cancelled invocation tasks did not exit before the cleanup deadline", null); + } + + private boolean awaitTasksUntil(long taskDeadline) throws InterruptedException { + while (true) { + var tasks = getActiveExecutorTasks(); + if (tasks.isEmpty()) { + return true; + } + var remainingNanos = remainingNanos(taskDeadline); + if (remainingNanos == 0) { + return false; + } + try { + CompletableFuture.allOf( + tasks.stream().map(ExecutorTaskHandle::exit).toArray(CompletableFuture[]::new)) + .get(remainingNanos, TimeUnit.NANOSECONDS); + } catch (ExecutionException impossible) { + throw new IllegalStateException("Executor task exit signal failed", impossible.getCause()); + } catch (TimeoutException timeout) { + return false; + } + } + } + + private void cancelActiveTasks() { + getActiveExecutorTasks().forEach(task -> task.cancel(true)); + } + + private long cleanupDeadlineNanos() { + var localDeadline = addToNow(MAX_CLEANUP_DURATION); + if (deadlineNanos == null) { + return localDeadline; + } + return Math.min(localDeadline, deadlineNanos - RESPONSE_HEADROOM.toNanos()); + } + + void beginDraining() { + synchronized (admissionLock) { + if (lifecycleState == LifecycleState.OPEN) { + lifecycleState = LifecycleState.DRAINING; } } } + private void closeLifecycle() { + synchronized (admissionLock) { + lifecycleState = LifecycleState.CLOSED; + } + } + + private T admitCheckpoint(Supplier request) { + synchronized (admissionLock) { + if (!checkpointAdmissionOpen || lifecycleState == LifecycleState.CLOSED) { + throw rejected("checkpoint request"); + } + return request.get(); + } + } + + private void stopCheckpointAdmission() { + synchronized (admissionLock) { + checkpointAdmissionOpen = false; + } + } + + private void requireOpen(String workType) { + if (lifecycleState != LifecycleState.OPEN) { + throw rejected(workType); + } + } + + private RejectedExecutionException rejected(String workType) { + return new RejectedExecutionException( + "Invocation " + invocationId + " is " + lifecycleState + "; cannot admit new " + workType); + } + /** Returns {@code true} if the given status represents a terminal (final) operation state. */ public static boolean isTerminalStatus(OperationStatus status) { return status == OperationStatus.SUCCEEDED @@ -488,6 +742,30 @@ private void stopAllOperations(Throwable cause) { registeredOperations.values().forEach(op -> op.getCompletionFuture().completeExceptionally(cause)); } + /** Waits for the root result, suspension, or invocation deadline on the Lambda runtime thread. */ + InvocationOutcome awaitInvocationOutcome(CompletableFuture userFuture) { + var outcomeFuture = runUntilCompleteOrSuspend(userFuture); + try { + var result = deadlineNanos == null + ? outcomeFuture.get() + : outcomeFuture.get( + remainingNanos(deadlineNanos - EXECUTION_DEADLINE_HEADROOM.toNanos()), + TimeUnit.NANOSECONDS); + return new InvocationOutcome<>(result, null); + } catch (InterruptedException interrupted) { + restoreInvocationThreadInterrupt = true; + var failure = lifecycleFailure("Invocation runtime thread was interrupted", interrupted); + signalLifecycleFailure(failure); + return new InvocationOutcome<>(null, failure); + } catch (TimeoutException timeout) { + var failure = lifecycleFailure("Invocation deadline reached before execution completed", timeout); + signalLifecycleFailure(failure); + return new InvocationOutcome<>(null, failure); + } catch (ExecutionException failure) { + return new InvocationOutcome<>(null, ExceptionHelper.unwrapCompletableFuture(failure.getCause())); + } + } + /** * return a future that completes when userFuture completes successfully or the execution is terminated or * suspended. @@ -506,4 +784,65 @@ public CompletableFuture runUntilCompleteOrSuspend(CompletableFuture u return null; }); } + + private static String resolveInvocationId(Context lambdaContext, String durableExecutionArn) { + if (lambdaContext != null) { + var requestId = lambdaContext.getAwsRequestId(); + if (requestId != null && !requestId.isBlank()) { + return requestId; + } + } + return durableExecutionArn + "#local-" + LOCAL_INVOCATION_SEQUENCE.incrementAndGet(); + } + + private static Long resolveDeadlineNanos(Context lambdaContext) { + if (lambdaContext == null) { + return null; + } + var remainingMillis = Math.max(0L, lambdaContext.getRemainingTimeInMillis()); + var remainingNanos = TimeUnit.MILLISECONDS.toNanos(remainingMillis); + var now = System.nanoTime(); + return now > Long.MAX_VALUE - remainingNanos ? Long.MAX_VALUE : now + remainingNanos; + } + + private void signalLifecycleFailure(Throwable failure) { + stopAllOperations(failure); + executionExceptionFuture.completeExceptionally(failure); + } + + private static UnrecoverableDurableExecutionException lifecycleFailure(String message, Throwable cause) { + return new UnrecoverableDurableExecutionException( + ErrorObject.builder() + .errorType(LIFECYCLE_ERROR_TYPE) + .errorMessage(message) + .build(), + true, + cause); + } + + private static Throwable combineFailures(Throwable primary, Throwable additional) { + if (primary == null) { + return additional; + } + primary.addSuppressed(additional); + return primary; + } + + private static long minDeadline(long deadline, Duration maximumWait) { + return Math.min(deadline, addToNow(maximumWait)); + } + + private static long addToNow(Duration duration) { + var now = System.nanoTime(); + var nanos = duration.toNanos(); + return now > Long.MAX_VALUE - nanos ? Long.MAX_VALUE : now + nanos; + } + + private static long remainingNanos(long deadline) { + return Math.max(0, deadline - System.nanoTime()); + } + + private static Duration remainingDuration(long deadline) { + return Duration.ofNanos(remainingNanos(deadline)); + } } diff --git a/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutorTaskHandle.java b/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutorTaskHandle.java new file mode 100644 index 000000000..8c9ee1fea --- /dev/null +++ b/sdk/src/main/java/software/amazon/lambda/durable/execution/ExecutorTaskHandle.java @@ -0,0 +1,182 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.execution; + +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.Future; +import java.util.concurrent.atomic.AtomicBoolean; +import java.util.concurrent.atomic.AtomicReference; +import java.util.function.Supplier; + +/** + * Tracks one executor task owned by a single Lambda invocation. + * + *

    This is runtime lifecycle metadata, not a checkpointed durable operation and not a request to invoke another + * Lambda function. Tasks include the root durable handler, user code for steps and child contexts, and SDK-owned + * map/parallel coordination. A durable operation may have no active task while it is replayed or waiting, and may + * create multiple tasks over its lifetime. + * + *

    The logical {@link #completion()} future reports the action's result to SDK code. The separate {@link #exit()} + * future proves that the executor wrapper has actually stopped running. {@link #execution()} retains the underlying + * executor handle so a later invocation-cleanup phase can request interruption without confusing logical completion + * with actual task exit. + */ +final class ExecutorTaskHandle { + + /** + * Runtime role used for invocation-local cleanup, diagnostics, and executor routing. This is deliberately separate + * from a persisted durable operation type: the root has no operation, and map/parallel coordinators share the + * CONTEXT operation type with child user code. + */ + enum Role { + /** The top-level durable handler submitted by {@link DurableExecutor}. */ + ROOT, + + /** User code for an operation whose type is STEP, including wait-for-condition checks. */ + STEP, + + /** User code running in a child context, map iteration, parallel branch, or retry context. */ + CHILD_CONTEXT, + + /** The SDK-owned scheduling loop for a map or parallel operation. */ + COORDINATOR + } + + enum State { + /** Admitted to the invocation scope but not yet started by the executor. */ + REGISTERED, + + /** The executor has entered the task wrapper and the action may still be running. */ + RUNNING, + + /** The wrapper has returned, or the executor accepted cancellation before the wrapper started. */ + EXITED + } + + /** Scope-local sequence number; this is not a durable operation ID. */ + private final long id; + + private final Role role; + + /** Associated durable operation ID, or {@code null} for the root handler. */ + private final String operationId; + + /** Removes this task from the execution manager's active-task registry. */ + private final Runnable onExit; + + /** Logical action result consumed by the SDK; cancellation can complete it before the action exits. */ + private final CompletableFuture completion = new CompletableFuture<>(); + + /** Completes only when no executor thread can still be running this task. */ + private final CompletableFuture exit = new CompletableFuture<>(); + + /** Handle returned by {@link java.util.concurrent.ExecutorService#submit(Runnable)}. */ + private final AtomicReference> execution = new AtomicReference<>(); + + private final AtomicReference state = new AtomicReference<>(State.REGISTERED); + private final AtomicBoolean cancellationRequested = new AtomicBoolean(); + private final AtomicBoolean interruptRequested = new AtomicBoolean(); + + ExecutorTaskHandle(long id, Role role, String operationId, Runnable onExit) { + this.id = id; + this.role = role; + this.operationId = operationId; + this.onExit = onExit; + } + + /** Executor wrapper that captures the logical result and always records actual exit. */ + void run(Supplier action) { + if (!state.compareAndSet(State.REGISTERED, State.RUNNING)) { + return; + } + try { + completion.complete(action.get()); + } catch (Throwable throwable) { + completion.completeExceptionally(throwable); + } finally { + markExited(); + } + } + + /** + * Binds the executor handle after submission. A cancellation racing with submission is remembered and applied as + * soon as this handle becomes available. + */ + void bindExecution(Future future) { + if (!execution.compareAndSet(null, future)) { + throw new IllegalStateException("Executor task already has an execution future"); + } + if (cancellationRequested.get()) { + cancelExecution(future); + } + } + + /** Records a synchronous executor rejection and releases the scope registration. */ + void submissionFailed(Throwable failure) { + completion.completeExceptionally(failure); + markExited(); + } + + /** + * Requests cooperative cancellation through the executor handle. Logical completion is cancelled immediately; + * actual exit remains pending until running code returns. User code that ignores interruption cannot be forcibly + * stopped. + */ + boolean cancel(boolean mayInterruptIfRunning) { + if (mayInterruptIfRunning) { + interruptRequested.set(true); + } + cancellationRequested.set(true); + completion.cancel(false); + var future = execution.get(); + if (future == null) { + return true; + } + return cancelExecution(future); + } + + private boolean cancelExecution(Future future) { + if (!future.cancel(interruptRequested.get())) { + return false; + } + if (state.compareAndSet(State.REGISTERED, State.EXITED)) { + onExit.run(); + exit.complete(null); + } + return true; + } + + private void markExited() { + state.set(State.EXITED); + onExit.run(); + exit.complete(null); + } + + long id() { + return id; + } + + Role role() { + return role; + } + + String operationId() { + return operationId; + } + + State state() { + return state.get(); + } + + CompletableFuture completion() { + return completion; + } + + CompletableFuture exit() { + return exit; + } + + Future execution() { + return execution.get(); + } +} diff --git a/sdk/src/main/java/software/amazon/lambda/durable/operation/BaseDurableOperation.java b/sdk/src/main/java/software/amazon/lambda/durable/operation/BaseDurableOperation.java index 5cd40820e..49a2423f9 100644 --- a/sdk/src/main/java/software/amazon/lambda/durable/operation/BaseDurableOperation.java +++ b/sdk/src/main/java/software/amazon/lambda/durable/operation/BaseDurableOperation.java @@ -329,8 +329,7 @@ protected void runUserHandler(Runnable runnable, ThreadType threadType) { // registerActiveThread is idempotent (no-op if already registered). registerActiveThread(operationId); - runningUserHandler.set(CompletableFuture.runAsync( - wrapped, getContext().getDurableConfig().getExecutorService())); + runningUserHandler.set(executionManager.submitOperationTask(this, wrapped)); } /** diff --git a/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java b/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java index 1aaba5d4c..dc60f6d90 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/TestUtils.java @@ -8,8 +8,11 @@ import java.util.ArrayList; import java.util.List; import java.util.UUID; +import java.util.concurrent.CompletableFuture; import software.amazon.awssdk.services.lambda.model.*; import software.amazon.lambda.durable.client.DurableExecutionClient; +import software.amazon.lambda.durable.context.DurableContextImpl; +import software.amazon.lambda.durable.execution.ExecutionManager; import software.amazon.lambda.durable.execution.OperationIdGenerator; public class TestUtils { @@ -69,4 +72,11 @@ public static DurableExecutionClient createMockClient() { public static String hashOperationId(String rawId) { return OperationIdGenerator.hashOperationId(rawId); } + + /** Makes a mocked ExecutionManager execute operation tasks like the real manager. */ + public static void executeOperationTasks(ExecutionManager manager, DurableContextImpl context) { + when(manager.submitOperationTask(any(), any())) + .thenAnswer(invocation -> CompletableFuture.runAsync( + invocation.getArgument(1), context.getDurableConfig().getExecutorService())); + } } diff --git a/sdk/src/test/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionExceptionTest.java b/sdk/src/test/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionExceptionTest.java index 67c2f0fa1..b3232e38a 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionExceptionTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/exception/UnrecoverableDurableExecutionExceptionTest.java @@ -5,6 +5,7 @@ import static org.junit.jupiter.api.Assertions.*; import org.junit.jupiter.api.Test; +import software.amazon.awssdk.services.lambda.model.ErrorObject; class UnrecoverableDurableExecutionExceptionTest { @@ -36,4 +37,15 @@ void testIllegalDurableOperationException() { assertInstanceOf(RuntimeException.class, exception); assertInstanceOf(DurableExecutionException.class, exception); } + + @Test + void preservesCauseForRetryableInvocationFailures() { + var cause = new InterruptedException("interrupted"); + var error = ErrorObject.builder().errorMessage("cleanup failed").build(); + + var exception = new UnrecoverableDurableExecutionException(error, true, cause); + + assertSame(cause, exception.getCause()); + assertTrue(exception.isRetryable()); + } } diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcherTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcherTest.java index 6aec8f7c1..355aa6412 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcherTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/ApiRequestDelayedBatcherTest.java @@ -18,6 +18,7 @@ import java.util.ArrayList; import java.util.List; import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeoutException; import java.util.function.Consumer; @@ -27,6 +28,7 @@ class ApiRequestDelayedBatcherTest { private static final Duration SHORT_DELAY = Duration.ofMillis(5); private static final Duration LONG_DELAY = Duration.ofMillis(100); + private static final Duration SHUTDOWN_TIMEOUT = Duration.ofSeconds(5); private static final int MAX_BATCH_SIZE = 3; private static final int MAX_BATCH_BINARY_SIZE_IN_BYTES = 200; @@ -191,11 +193,11 @@ void whenBatchActionThrowsRuntimeException_futuresReceiveOriginalCause() { } @Test - void whenShutdownCalled_pendingItemsAreFlushedImmediately() { + void whenShutdownCalled_pendingItemsAreFlushedImmediately() throws Exception { var future = cut.submit(input, LONG_DELAY); assertFalse(future.isDone()); - cut.shutdown(); + cut.shutdown(SHUTDOWN_TIMEOUT); assertTrue(future.isDone()); verify(doBatchAction).accept(any()); @@ -211,19 +213,57 @@ void whenEarlierDelaySubmitted_batchFlushesAtEarlierTime() { } @Test - void whenNoItemsSubmitted_shutdownDoesNotInvokeBatchAction() { - cut.shutdown(); + void whenDelayExceedsNinetyMinutes_firstSubmissionStillSetsFlushDeadline() throws Exception { + var submittedAt = System.nanoTime(); + cut.submit(input, Duration.ofMinutes(91)); + + var flushTimeField = ApiRequestDelayedBatcher.class.getDeclaredField("delayedBatchFlushTime"); + flushTimeField.setAccessible(true); + var flushTime = flushTimeField.getLong(cut); + + assertTrue(flushTime > submittedAt + Duration.ofMinutes(90).toNanos()); + cut.shutdown(SHUTDOWN_TIMEOUT); + } + + @Test + void whenNoItemsSubmitted_shutdownDoesNotInvokeBatchAction() throws Exception { + cut.shutdown(SHUTDOWN_TIMEOUT); verify(doBatchAction, never()).accept(any()); } @Test - void whenMultipleBatchesFlushedViaShutdown_allFuturesComplete() { + void whenMultipleBatchesFlushedViaShutdown_allFuturesComplete() throws Exception { var future1 = cut.submit(input, LONG_DELAY); - cut.shutdown(); + cut.shutdown(SHUTDOWN_TIMEOUT); assertTrue(future1.isDone()); var future2 = cut.submit(input, LONG_DELAY); - cut.shutdown(); + cut.shutdown(SHUTDOWN_TIMEOUT); assertTrue(future2.isDone()); } + + @Test + void shutdownHonorsTimeoutWhileBatchActionIsBlocked() throws Exception { + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var batcher = new ApiRequestDelayedBatcher( + MAX_BATCH_SIZE, MAX_BATCH_BINARY_SIZE_IN_BYTES, item -> 0, items -> { + entered.countDown(); + try { + release.await(); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + } + }); + var future = batcher.submit(input, Duration.ZERO); + assertTrue(entered.await(5, TimeUnit.SECONDS)); + + try { + assertThrows(TimeoutException.class, () -> batcher.shutdown(Duration.ofMillis(50))); + assertFalse(future.isDone()); + } finally { + release.countDown(); + } + future.get(5, TimeUnit.SECONDS); + } } diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/CheckpointManagerTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/CheckpointManagerTest.java index 61f534e26..e9b609147 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/execution/CheckpointManagerTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/CheckpointManagerTest.java @@ -34,6 +34,8 @@ class CheckpointManagerTest { + private static final Duration SHUTDOWN_TIMEOUT = Duration.ofSeconds(5); + private DurableConfig config; private DurableExecutionClient client; private CheckpointManager batcher; @@ -171,11 +173,11 @@ void pollForUpdate_handlesMultiplePollers() throws Exception { } @Test - void shutdown_completesAllPendingPollersWithException() { + void shutdown_completesAllPendingPollersWithException() throws Exception { var future1 = batcher.pollForUpdate("op-1"); var future2 = batcher.pollForUpdate("op-2"); - batcher.shutdown(); + batcher.shutdown(SHUTDOWN_TIMEOUT); assertTrue(future1.isCompletedExceptionally()); assertTrue(future2.isCompletedExceptionally()); @@ -197,7 +199,7 @@ void shutdown_waitsForPendingCheckpoints() throws Exception { .type(OperationType.STEP) .build()); - batcher.shutdown(); + batcher.shutdown(SHUTDOWN_TIMEOUT); assertTrue(future.isDone()); verify(client, atLeastOnce()).checkpoint(anyString(), anyString(), anyList()); diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerLifecycleTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerLifecycleTest.java new file mode 100644 index 000000000..edd55e7ec --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerLifecycleTest.java @@ -0,0 +1,400 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.execution; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertInstanceOf; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertNotNull; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; +import static software.amazon.lambda.durable.model.ExecutionStatus.PENDING; + +import com.amazonaws.services.lambda.runtime.Context; +import java.time.Duration; +import java.time.Instant; +import java.util.List; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.FutureTask; +import java.util.concurrent.RejectedExecutionException; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.concurrent.atomic.AtomicInteger; +import org.junit.jupiter.api.Test; +import software.amazon.awssdk.services.lambda.model.CheckpointUpdatedExecutionState; +import software.amazon.awssdk.services.lambda.model.ExecutionDetails; +import software.amazon.awssdk.services.lambda.model.Operation; +import software.amazon.awssdk.services.lambda.model.OperationAction; +import software.amazon.awssdk.services.lambda.model.OperationStatus; +import software.amazon.awssdk.services.lambda.model.OperationType; +import software.amazon.awssdk.services.lambda.model.OperationUpdate; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TestUtils; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.exception.UnrecoverableDurableExecutionException; +import software.amazon.lambda.durable.model.DurableExecutionInput; +import software.amazon.lambda.durable.operation.BaseDurableOperation; +import software.amazon.lambda.durable.plugin.DurableExecutionPlugin; +import software.amazon.lambda.durable.plugin.InvocationEndInfo; + +class ExecutionManagerLifecycleTest { + private static final String EXECUTION_ID = "execution"; + private static final String EXECUTION_ARN = + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/test/" + EXECUTION_ID; + + @Test + void capturesRequestIdAndInvocationDeadline() { + var context = mock(Context.class); + when(context.getAwsRequestId()).thenReturn("request-id"); + when(context.getRemainingTimeInMillis()).thenReturn(2_000); + var manager = createManager(context, null); + + try { + assertEquals("request-id", manager.getInvocationId()); + var remaining = manager.getRemainingInvocationTime().orElseThrow(); + assertTrue(remaining.compareTo(Duration.ZERO) > 0); + assertTrue(remaining.compareTo(Duration.ofSeconds(2)) <= 0); + } finally { + manager.close(); + } + } + + @Test + void createsDistinctLocalIdsWithoutInventingADeadline() { + var first = createManager(null, null); + var second = createManager(null, null); + + try { + assertNotEquals(first.getInvocationId(), second.getInvocationId()); + assertTrue(first.getRemainingInvocationTime().isEmpty()); + } finally { + first.close(); + second.close(); + } + } + + @Test + void registersTaskBeforeQueuedExecutionAndTracksItsExit() throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var blockerEntered = new CountDownLatch(1); + var releaseBlocker = new CountDownLatch(1); + executor.submit(() -> { + blockerEntered.countDown(); + await(releaseBlocker); + }); + assertTrue(blockerEntered.await(5, TimeUnit.SECONDS)); + var manager = createManager(null, executor); + + try { + var completion = manager.submitRootTask(() -> "result"); + var task = manager.getActiveExecutorTasks().get(0); + + assertEquals(ExecutorTaskHandle.State.REGISTERED, task.state()); + assertNotNull(task.execution()); + assertFalse(completion.isDone()); + assertFalse(task.exit().isDone()); + + releaseBlocker.countDown(); + assertEquals("result", completion.get(5, TimeUnit.SECONDS)); + task.exit().get(5, TimeUnit.SECONDS); + assertEquals(ExecutorTaskHandle.State.EXITED, task.state()); + assertTrue(manager.getActiveExecutorTasks().isEmpty()); + } finally { + releaseBlocker.countDown(); + manager.close(); + stop(executor); + } + } + + @Test + void taskCancellationSeparatesLogicalCompletionFromActualExit() throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var entered = new CountDownLatch(1); + var interrupted = new CountDownLatch(1); + var release = new CountDownLatch(1); + var manager = createManager(null, executor); + try { + var completion = manager.submitRootTask(() -> { + entered.countDown(); + while (release.getCount() > 0) { + try { + release.await(); + } catch (InterruptedException expected) { + interrupted.countDown(); + } + } + return "ignored"; + }); + assertTrue(entered.await(5, TimeUnit.SECONDS)); + var task = manager.getActiveExecutorTasks().get(0); + + assertTrue(task.cancel(true)); + assertTrue(interrupted.await(5, TimeUnit.SECONDS)); + assertTrue(completion.isCancelled()); + assertFalse(task.exit().isDone()); + + release.countDown(); + task.exit().get(5, TimeUnit.SECONDS); + assertEquals(ExecutorTaskHandle.State.EXITED, task.state()); + } finally { + release.countDown(); + manager.close(); + stop(executor); + } + } + + @Test + void cancellationRequestedBeforeExecutorHandleIsBoundIsNotLost() throws Exception { + var exits = new AtomicInteger(); + var task = new ExecutorTaskHandle(1, ExecutorTaskHandle.Role.ROOT, null, exits::incrementAndGet); + var execution = new FutureTask(() -> null); + + assertTrue(task.cancel(true)); + assertTrue(task.completion().isCancelled()); + assertFalse(task.exit().isDone()); + + task.bindExecution(execution); + + assertTrue(execution.isCancelled()); + task.exit().get(5, TimeUnit.SECONDS); + assertEquals(1, exits.get()); + } + + @Test + void failedSubmissionIsRemovedFromManager() { + var executor = Executors.newSingleThreadExecutor(); + executor.shutdown(); + var manager = createManager(null, executor); + + try { + assertThrows(RejectedExecutionException.class, () -> manager.submitRootTask(() -> "result")); + assertTrue(manager.getActiveExecutorTasks().isEmpty()); + } finally { + manager.close(); + } + } + + @Test + void drainingRejectsNewWorkButAllowsExistingCheckpointCleanup() throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var manager = createManager(null, executor); + manager.beginDraining(); + + try { + assertEquals(ExecutionManager.LifecycleState.DRAINING, manager.getLifecycleState()); + assertThrows( + RejectedExecutionException.class, + () -> manager.registerOperation(mock(BaseDurableOperation.class))); + assertThrows(RejectedExecutionException.class, () -> manager.submitRootTask(() -> "result")); + + var checkpoint = manager.sendOperationUpdate(OperationUpdate.builder() + .id("step") + .name("step") + .type(OperationType.STEP) + .subType("Step") + .action(OperationAction.START) + .build()); + manager.close(); + checkpoint.get(5, TimeUnit.SECONDS); + + assertEquals(ExecutionManager.LifecycleState.CLOSED, manager.getLifecycleState()); + assertThrows( + RejectedExecutionException.class, + () -> manager.sendOperationUpdate(OperationUpdate.builder().build())); + } finally { + stop(executor); + } + } + + @Test + void interruptedDrainStillShutsDownCheckpointsAndRestoresInterrupt() throws Exception { + var executor = Executors.newSingleThreadExecutor(); + var taskEntered = new CountDownLatch(1); + var taskInterrupted = new CountDownLatch(1); + var manager = createManager(null, executor); + manager.submitRootTask(() -> { + taskEntered.countDown(); + try { + new CountDownLatch(1).await(); + } catch (InterruptedException expected) { + taskInterrupted.countDown(); + } + return "cancelled"; + }); + assertTrue(taskEntered.await(5, TimeUnit.SECONDS)); + var checkpoint = manager.sendOperationUpdate(OperationUpdate.builder() + .id("step") + .name("step") + .type(OperationType.STEP) + .subType("Step") + .action(OperationAction.START) + .build()); + + Thread.currentThread().interrupt(); + try { + var failure = manager.finishInvocation(null); + + var lifecycleFailure = assertInstanceOf(UnrecoverableDurableExecutionException.class, failure); + assertInstanceOf(InterruptedException.class, lifecycleFailure.getCause()); + assertEquals(ExecutionManager.LifecycleState.CLOSED, manager.getLifecycleState()); + assertTrue(checkpoint.isDone(), "checkpoint shutdown should still be attempted"); + assertTrue(Thread.currentThread().isInterrupted()); + Thread.interrupted(); + assertTrue(taskInterrupted.await(5, TimeUnit.SECONDS)); + } finally { + Thread.interrupted(); + stop(executor); + } + } + + @Test + void invocationDeadlineCancelsAStuckRootAndReturnsRetryableFailure() throws Exception { + var context = mock(Context.class); + when(context.getAwsRequestId()).thenReturn("deadline-request"); + when(context.getRemainingTimeInMillis()).thenReturn(700); + var executor = Executors.newSingleThreadExecutor(); + var entered = new CountDownLatch(1); + var interrupted = new CountDownLatch(1); + var manager = createManager(context, executor); + var root = manager.submitRootTask(() -> { + entered.countDown(); + try { + new CountDownLatch(1).await(); + } catch (InterruptedException expected) { + interrupted.countDown(); + } + return "cancelled"; + }); + assertTrue(entered.await(5, TimeUnit.SECONDS)); + + try { + var outcome = manager.awaitInvocationOutcome(root); + var deadlineFailure = assertInstanceOf(UnrecoverableDurableExecutionException.class, outcome.failure()); + assertTrue(deadlineFailure.isRetryable()); + + assertEquals(deadlineFailure, manager.finishInvocation(deadlineFailure)); + assertTrue(interrupted.await(5, TimeUnit.SECONDS)); + assertEquals(ExecutionManager.LifecycleState.CLOSED, manager.getLifecycleState()); + } finally { + stop(executor); + } + } + + @Test + void responseCriticalCheckpointWaitIsBoundedByInvocationDeadline() { + var context = mock(Context.class); + when(context.getAwsRequestId()).thenReturn("checkpoint-deadline-request"); + when(context.getRemainingTimeInMillis()).thenReturn(250); + var manager = createManager(context, null); + + var failure = assertThrows( + UnrecoverableDurableExecutionException.class, + () -> manager.awaitCheckpointCompletion(new CompletableFuture<>())); + + assertTrue(failure.isRetryable()); + assertInstanceOf(TimeoutException.class, failure.getCause()); + assertEquals(failure, manager.finishInvocation(failure)); + } + + @Test + void pendingResponseAndInvocationEndWaitForRootFinally() throws Exception { + var users = Executors.newCachedThreadPool(); + var runtime = Executors.newSingleThreadExecutor(); + var cleanupEntered = new CountDownLatch(1); + var releaseCleanup = new CountDownLatch(1); + var invocationEnded = new CountDownLatch(1); + var config = DurableConfig.builder() + .withDurableExecutionClient(TestUtils.createMockClient()) + .withExecutorService(users) + .withPlugins(new DurableExecutionPlugin() { + @Override + public void onInvocationEnd(InvocationEndInfo info) { + invocationEnded.countDown(); + } + }) + .build(); + + try { + var response = runtime.submit(() -> DurableExecutor.execute( + input(), + context(10_000), + TypeToken.get(String.class), + (value, durableContext) -> { + try { + durableContext.wait("wait", Duration.ofSeconds(5)); + return value; + } finally { + cleanupEntered.countDown(); + await(releaseCleanup); + } + }, + config)); + assertTrue(cleanupEntered.await(5, TimeUnit.SECONDS)); + + assertThrows(TimeoutException.class, () -> response.get(100, TimeUnit.MILLISECONDS)); + assertEquals(1, invocationEnded.getCount()); + + releaseCleanup.countDown(); + assertEquals(PENDING, response.get(5, TimeUnit.SECONDS).status()); + assertTrue(invocationEnded.await(5, TimeUnit.SECONDS)); + } finally { + releaseCleanup.countDown(); + stop(users); + stop(runtime); + } + } + + private ExecutionManager createManager(Context context, ExecutorService executor) { + var config = DurableConfig.builder().withDurableExecutionClient(TestUtils.createMockClient()); + if (executor != null) { + config.withExecutorService(executor); + } + return new ExecutionManager(input(), config.build(), context); + } + + private DurableExecutionInput input() { + var execution = Operation.builder() + .id(EXECUTION_ID) + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .startTimestamp(Instant.now()) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + return new DurableExecutionInput( + EXECUTION_ARN, + "token", + CheckpointUpdatedExecutionState.builder() + .operations(List.of(execution)) + .build()); + } + + private Context context(int remainingMillis) { + var context = mock(Context.class); + when(context.getAwsRequestId()).thenReturn("request-id"); + when(context.getRemainingTimeInMillis()).thenReturn(remainingMillis); + return context; + } + + private static void await(CountDownLatch latch) { + try { + if (!latch.await(5, TimeUnit.SECONDS)) { + throw new IllegalStateException("Test synchronization timed out"); + } + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new IllegalStateException(interrupted); + } + } + + private static void stop(ExecutorService executor) throws InterruptedException { + executor.shutdownNow(); + assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS)); + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerTest.java index 056c80e1a..dd99ea070 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/ExecutionManagerTest.java @@ -315,4 +315,67 @@ protected void replay(Operation existing) { operation.execute(); assertTrue(operation.getCompletionFuture().isDone()); } + + @Test + void tracksRootStepChildAndCoordinatorTasks() throws Exception { + var manager = createManager(List.of(executionOp())); + var entered = new CountDownLatch(5); + var release = new CountDownLatch(1); + try { + var root = manager.submitRootTask(() -> { + waitForRelease(entered, release); + return "root-result"; + }); + var step = manager.submitOperationTask( + operation("step", OperationSubType.STEP), () -> waitForRelease(entered, release)); + var waitForCondition = manager.submitOperationTask( + operation("wait", OperationSubType.WAIT_FOR_CONDITION), () -> waitForRelease(entered, release)); + var child = manager.submitOperationTask( + operation("child", OperationSubType.RUN_IN_CHILD_CONTEXT), () -> waitForRelease(entered, release)); + var coordinator = manager.submitOperationTask( + operation("map", OperationSubType.MAP), () -> waitForRelease(entered, release)); + + assertTrue(entered.await(5, TimeUnit.SECONDS)); + + var tasks = manager.getActiveExecutorTasks(); + assertEquals( + List.of( + ExecutorTaskHandle.Role.ROOT, + ExecutorTaskHandle.Role.STEP, + ExecutorTaskHandle.Role.STEP, + ExecutorTaskHandle.Role.CHILD_CONTEXT, + ExecutorTaskHandle.Role.COORDINATOR), + tasks.stream().map(ExecutorTaskHandle::role).toList()); + assertTrue(tasks.stream().allMatch(task -> task.state() == ExecutorTaskHandle.State.RUNNING)); + + release.countDown(); + assertEquals("root-result", root.get(5, TimeUnit.SECONDS)); + step.get(5, TimeUnit.SECONDS); + waitForCondition.get(5, TimeUnit.SECONDS); + child.get(5, TimeUnit.SECONDS); + coordinator.get(5, TimeUnit.SECONDS); + assertTrue(manager.getActiveExecutorTasks().isEmpty()); + } finally { + release.countDown(); + manager.close(); + } + } + + private void waitForRelease(CountDownLatch entered, CountDownLatch release) { + entered.countDown(); + try { + assertTrue(release.await(5, TimeUnit.SECONDS)); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new IllegalStateException(interrupted); + } + } + + private BaseDurableOperation operation(String id, OperationSubType subType) { + var operation = mock(BaseDurableOperation.class); + when(operation.getOperationId()).thenReturn(id); + when(operation.getSubType()).thenReturn(subType); + when(operation.getType()).thenReturn(subType.getOperationType()); + return operation; + } } diff --git a/sdk/src/test/java/software/amazon/lambda/durable/execution/LmiLifecycleRegressionTest.java b/sdk/src/test/java/software/amazon/lambda/durable/execution/LmiLifecycleRegressionTest.java new file mode 100644 index 000000000..6d53fc018 --- /dev/null +++ b/sdk/src/test/java/software/amazon/lambda/durable/execution/LmiLifecycleRegressionTest.java @@ -0,0 +1,198 @@ +// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +// SPDX-License-Identifier: Apache-2.0 +package software.amazon.lambda.durable.execution; + +import static org.junit.jupiter.api.Assertions.*; +import static org.mockito.Mockito.*; +import static software.amazon.lambda.durable.model.ExecutionStatus.PENDING; +import static software.amazon.lambda.durable.model.ExecutionStatus.SUCCEEDED; + +import com.amazonaws.services.lambda.runtime.Context; +import java.time.Duration; +import java.util.List; +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.ThreadPoolExecutor; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.TimeoutException; +import java.util.stream.IntStream; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.condition.EnabledIfSystemProperty; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; +import software.amazon.awssdk.services.lambda.model.CheckpointUpdatedExecutionState; +import software.amazon.awssdk.services.lambda.model.ExecutionDetails; +import software.amazon.awssdk.services.lambda.model.Operation; +import software.amazon.awssdk.services.lambda.model.OperationStatus; +import software.amazon.awssdk.services.lambda.model.OperationType; +import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TestUtils; +import software.amazon.lambda.durable.TypeToken; +import software.amazon.lambda.durable.model.DurableExecutionInput; + +/** Local red regressions for #726. Real LMI validation lives in lmi-tests; these do not claim cloud coverage. */ +@EnabledIfSystemProperty(named = "test.lmi.regressions.enabled", matches = "true") +class LmiLifecycleRegressionTest { + @Test + void pendingMustWaitForRootFinally() throws Exception { + var cleanupEntered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var users = Executors.newCachedThreadPool(); + var runtime = Executors.newSingleThreadExecutor(); + try { + var response = runtime.submit(() -> DurableExecutor.execute( + input("pending"), + context(10_000), + TypeToken.get(String.class), + (value, ctx) -> { + try { + ctx.wait("wait", Duration.ofSeconds(5)); + return value; + } finally { + cleanupEntered.countDown(); + await(release); + } + }, + config(users))); + await(cleanupEntered); + assertThrows( + TimeoutException.class, + () -> response.get(250, TimeUnit.MILLISECONDS), + "PENDING must not precede the root handler's finally exit"); + release.countDown(); + assertEquals(PENDING, response.get(3, TimeUnit.SECONDS).status()); + } finally { + release.countDown(); + stop(users, runtime); + } + } + + @Test + void deadlineMustCancelAnInvocationOwnedTaskAndBoundCleanup() throws Exception { + var entered = new CountDownLatch(1); + var release = new CountDownLatch(1); + var exited = new CountDownLatch(1); + var users = Executors.newCachedThreadPool(); + var runtime = Executors.newSingleThreadExecutor(); + try { + runtime.submit(() -> DurableExecutor.execute( + input("deadline"), + context(1000), + TypeToken.get(String.class), + (value, ctx) -> { + ctx.stepAsync("blocked", String.class, step -> { + entered.countDown(); + try { + await(release); + return value; + } finally { + exited.countDown(); + } + }); + await(entered); + return value; + }, + config(users))); + await(entered); + assertTrue(exited.await(3, TimeUnit.SECONDS), "Task outlived the invocation deadline without cancellation"); + } finally { + release.countDown(); + stop(users, runtime); + } + } + + @ParameterizedTest + @ValueSource(booleans = {false, true}) + void sharedFixedExecutorMustProgress(boolean nested) throws Exception { + var entered = new CountDownLatch(2); + var users = (ThreadPoolExecutor) Executors.newFixedThreadPool(2); + var runtime = Executors.newFixedThreadPool(2); + try { + var config = config(users); + var responses = IntStream.range(0, 2) + .mapToObj(index -> runtime.submit(() -> DurableExecutor.execute( + input("fixed-" + index), + context(10_000), + TypeToken.get(String.class), + (value, ctx) -> { + entered.countDown(); + await(entered); + if (nested) { + return ctx.runInChildContext( + "child", + String.class, + child -> child.runInChildContext( + "grandchild", + String.class, + grandchild -> + grandchild.step("step", String.class, step -> value))); + } + return ctx.step("step", String.class, step -> value); + }, + config))) + .toList(); + await(entered); + for (var response : responses) { + assertEquals( + SUCCEEDED, + assertDoesNotThrow( + () -> response.get(2, TimeUnit.SECONDS), + "Shared executor progress deadline exceeded") + .status(), + "Blocking orchestration must not starve the work it awaits"); + } + } finally { + users.setMaximumPoolSize(16); + users.setCorePoolSize(16); // Escape only after the assertion, so the test process cannot hang. + stop(users, runtime); + } + } + + private static DurableConfig config(ExecutorService users) { + return DurableConfig.builder() + .withDurableExecutionClient(TestUtils.createMockClient()) + .withExecutorService(users) + .build(); + } + + private static Context context(int milliseconds) { + var context = mock(Context.class); + var deadline = System.nanoTime() + milliseconds * 1_000_000L; + when(context.getAwsRequestId()).thenReturn("test-request"); + when(context.getRemainingTimeInMillis()).thenAnswer(call -> (int) ((deadline - System.nanoTime()) / 1_000_000)); + return context; + } + + private static DurableExecutionInput input(String name) { + var operation = Operation.builder() + .id(name) + .type(OperationType.EXECUTION) + .status(OperationStatus.STARTED) + .executionDetails( + ExecutionDetails.builder().inputPayload("\"input\"").build()) + .build(); + return new DurableExecutionInput( + "arn:aws:lambda:us-east-1:123456789012:function:test/durable-execution/" + name, + "token", + CheckpointUpdatedExecutionState.builder() + .operations(List.of(operation)) + .build()); + } + + private static void await(CountDownLatch latch) { + try { + assertTrue(latch.await(8, TimeUnit.SECONDS), "Fixture coordination exceeded its bound"); + } catch (InterruptedException interrupted) { + Thread.currentThread().interrupt(); + throw new IllegalStateException(interrupted); + } + } + + private static void stop(ExecutorService users, ExecutorService runtime) throws InterruptedException { + runtime.shutdown(); + assertTrue(runtime.awaitTermination(10, TimeUnit.SECONDS)); + users.shutdown(); + assertTrue(users.awaitTermination(10, TimeUnit.SECONDS)); + } +} diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/ChildContextOperationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/ChildContextOperationTest.java index 99d994538..415d38028 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/operation/ChildContextOperationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/ChildContextOperationTest.java @@ -19,6 +19,7 @@ import software.amazon.awssdk.services.lambda.model.OperationType; import software.amazon.lambda.durable.DurableConfig; import software.amazon.lambda.durable.DurableContext; +import software.amazon.lambda.durable.TestUtils; import software.amazon.lambda.durable.TypeToken; import software.amazon.lambda.durable.config.RunInChildContextConfig; import software.amazon.lambda.durable.context.DurableContextImpl; @@ -73,6 +74,7 @@ void setUp() { when(durableContext.getExecutionManager()).thenReturn(executionManager); when(executionManager.getCurrentThreadContext()).thenReturn(new ThreadContext("Root", ThreadType.CONTEXT)); when(durableContext.getDurableConfig()).thenReturn(createConfig()); + TestUtils.executeOperationTasks(executionManager, durableContext); } private DurableConfig createConfig() { diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/ConcurrencyOperationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/ConcurrencyOperationTest.java index b6488139f..e713f66d4 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/operation/ConcurrencyOperationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/ConcurrencyOperationTest.java @@ -90,6 +90,7 @@ void setUp() { .status(OperationStatus.SUCCEEDED) .build()); when(executionManager.sendOperationUpdate(any())).thenReturn(CompletableFuture.completedFuture(null)); + TestUtils.executeOperationTasks(executionManager, durableContext); } private TestConcurrencyOperation createOperation(CompletionConfig completionConfig) throws Exception { diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/ParallelOperationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/ParallelOperationTest.java index e02287cdc..5dfc97758 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/operation/ParallelOperationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/ParallelOperationTest.java @@ -80,6 +80,7 @@ void setUp() { .build()); when(durableContext.createChildContext(anyString(), anyString(), eq(false))) .thenReturn(childContext); + TestUtils.executeOperationTasks(executionManager, durableContext); // Capture registered operations so we can drive onCheckpointComplete callbacks. var registeredOps = new ConcurrentHashMap(); diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/StepOperationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/StepOperationTest.java index be4962d71..2d87c12e9 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/operation/StepOperationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/StepOperationTest.java @@ -14,6 +14,7 @@ import software.amazon.awssdk.services.lambda.model.OperationStatus; import software.amazon.awssdk.services.lambda.model.StepDetails; import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TestUtils; import software.amazon.lambda.durable.TypeToken; import software.amazon.lambda.durable.config.StepConfig; import software.amazon.lambda.durable.context.DurableContextImpl; @@ -46,6 +47,7 @@ void setUp() { .thenReturn(DurableConfig.builder() .withExecutorService(Executors.newCachedThreadPool()) .build()); + TestUtils.executeOperationTasks(executionManager, durableContext); } private void mockFailedOperation( diff --git a/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionOperationTest.java b/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionOperationTest.java index c33f0160d..2739bad87 100644 --- a/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionOperationTest.java +++ b/sdk/src/test/java/software/amazon/lambda/durable/operation/WaitForConditionOperationTest.java @@ -19,6 +19,7 @@ import software.amazon.awssdk.services.lambda.model.OperationType; import software.amazon.awssdk.services.lambda.model.StepDetails; import software.amazon.lambda.durable.DurableConfig; +import software.amazon.lambda.durable.TestUtils; import software.amazon.lambda.durable.TypeToken; import software.amazon.lambda.durable.config.WaitForConditionConfig; import software.amazon.lambda.durable.context.DurableContextImpl; @@ -54,6 +55,7 @@ void setUp() { .thenReturn(DurableConfig.builder() .withExecutorService(Executors.newCachedThreadPool()) .build()); + TestUtils.executeOperationTasks(executionManager, durableContext); } private WaitForConditionOperation createOperation(